diff --git a/api/changelog.d/lapsed-invitations-block-re-invites.fixed.md b/api/changelog.d/lapsed-invitations-block-re-invites.fixed.md new file mode 100644 index 0000000000..29a8e99235 --- /dev/null +++ b/api/changelog.d/lapsed-invitations-block-re-invites.fixed.md @@ -0,0 +1 @@ +Lapsed pending invitations are reported as expired and no longer block a new invitation for the same email diff --git a/api/src/backend/api/filters.py b/api/src/backend/api/filters.py index a7f888b453..ad89671c89 100644 --- a/api/src/backend/api/filters.py +++ b/api/src/backend/api/filters.py @@ -1439,8 +1439,20 @@ 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") + state = ChoiceFilter(choices=Invitation.State.choices, method="filter_state") + state__in = ChoiceInFilter( + choices=Invitation.State.choices, lookup_expr="in", method="filter_state_in" + ) + + def filter_state(self, queryset, name, value): + return self.filter_state_in(queryset, name, [value]) + + def filter_state_in(self, queryset, name, value): + lapsed = Invitation.lapsed_q() + query = Q(state__in=value) & ~lapsed + if Invitation.State.EXPIRED in value: + query |= lapsed + return queryset.filter(query) class Meta: model = Invitation diff --git a/api/src/backend/api/models.py b/api/src/backend/api/models.py index a280708d53..6a18cd1994 100644 --- a/api/src/backend/api/models.py +++ b/api/src/backend/api/models.py @@ -1380,6 +1380,15 @@ class Invitation(RowLevelSecurityProtectedModel): self.email = self.email.strip().lower() super().save(*args, **kwargs) + @classmethod + def lapsed_q(cls): + """Pending invitations whose expiry date has already passed.""" + return Q(state=cls.State.PENDING, expires_at__lte=datetime.now(UTC)) + + @property + def is_lapsed(self): + return self.state == self.State.PENDING and self.expires_at <= datetime.now(UTC) + class Meta(RowLevelSecurityProtectedModel.Meta): db_table = "invitations" diff --git a/api/src/backend/api/tests/test_views.py b/api/src/backend/api/tests/test_views.py index a881a7becf..37b4a87729 100644 --- a/api/src/backend/api/tests/test_views.py +++ b/api/src/backend/api/tests/test_views.py @@ -8784,6 +8784,190 @@ class TestInvitationViewSet: user.id ) + @staticmethod + def _invitation_create_payload(email, role): + return json.dumps( + { + "data": { + "type": "invitations", + "attributes": {"email": email}, + "relationships": { + "roles": {"data": [{"type": "roles", "id": str(role.id)}]} + }, + } + } + ) + + @staticmethod + def _create_lapsed_invitation(email, tenant, inviter): + return Invitation.objects.create( + email=email, + state=Invitation.State.PENDING, + expires_at=datetime.now(UTC) - timedelta(days=1), + inviter=inviter, + tenant=tenant, + ) + + def test_invitations_create_with_lapsed_pending_invitation_for_same_email( + self, + authenticated_client, + create_test_user, + tenants_fixture, + invitations_fixture, + roles_fixture, + ): + lapsed_invitation, expired_invitation = invitations_fixture + lapsed_invitation.expires_at = datetime.now(UTC) - timedelta(days=1) + lapsed_invitation.save() + other_email_lapsed_invitation = self._create_lapsed_invitation( + "other@prowler.com", tenants_fixture[0], create_test_user + ) + + response = authenticated_client.post( + reverse("invitation-list"), + data=self._invitation_create_payload( + lapsed_invitation.email, roles_fixture[0] + ), + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_201_CREATED + new_invitation = Invitation.objects.get(id=response.json()["data"]["id"]) + assert new_invitation.email == lapsed_invitation.email + assert new_invitation.state == Invitation.State.PENDING + lapsed_invitation.refresh_from_db() + assert lapsed_invitation.state == Invitation.State.EXPIRED + expired_invitation.refresh_from_db() + assert expired_invitation.state == Invitation.State.EXPIRED + other_email_lapsed_invitation.refresh_from_db() + assert other_email_lapsed_invitation.state == Invitation.State.PENDING + + def test_invitations_create_with_active_pending_invitation_for_same_email( + self, + authenticated_client, + create_test_user, + tenants_fixture, + invitations_fixture, + roles_fixture, + ): + active_invitation, _ = invitations_fixture + self._create_lapsed_invitation( + active_invitation.email, tenants_fixture[0], create_test_user + ) + invitation_count = Invitation.objects.count() + + response = authenticated_client.post( + reverse("invitation-list"), + data=self._invitation_create_payload( + active_invitation.email, roles_fixture[0] + ), + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_400_BAD_REQUEST + assert ( + response.json()["errors"][0]["source"]["pointer"] + == "/data/attributes/email" + ) + assert Invitation.objects.count() == invitation_count + active_invitation.refresh_from_db() + assert active_invitation.state == Invitation.State.PENDING + + def test_invitations_create_ignores_pending_invitations_from_other_tenants( + self, authenticated_client, create_test_user, tenants_fixture, roles_fixture + ): + email = "cross_tenant@prowler.com" + other_tenant = tenants_fixture[1] + other_tenant_lapsed_invitation = self._create_lapsed_invitation( + email, other_tenant, create_test_user + ) + Invitation.objects.create( + email=email, inviter=create_test_user, tenant=other_tenant + ) + + response = authenticated_client.post( + reverse("invitation-list"), + data=self._invitation_create_payload(email, roles_fixture[0]), + content_type="application/vnd.api+json", + ) + + assert response.status_code == status.HTTP_201_CREATED + other_tenant_lapsed_invitation.refresh_from_db() + assert other_tenant_lapsed_invitation.state == Invitation.State.PENDING + + def test_invitations_report_lapsed_pending_invitation_as_expired( + self, + authenticated_client, + create_test_user, + tenants_fixture, + invitations_fixture, + ): + active_invitation, expired_invitation = invitations_fixture + lapsed_invitation = self._create_lapsed_invitation( + "lapsed@prowler.com", tenants_fixture[0], create_test_user + ) + + list_response = authenticated_client.get(reverse("invitation-list")) + retrieve_response = authenticated_client.get( + reverse("invitation-detail", kwargs={"pk": lapsed_invitation.id}) + ) + + assert list_response.status_code == status.HTTP_200_OK + assert retrieve_response.status_code == status.HTTP_200_OK + assert { + invitation["id"]: invitation["attributes"]["state"] + for invitation in list_response.json()["data"] + } == { + str(active_invitation.id): Invitation.State.PENDING.value, + str(expired_invitation.id): Invitation.State.EXPIRED.value, + str(lapsed_invitation.id): Invitation.State.EXPIRED.value, + } + assert ( + retrieve_response.json()["data"]["attributes"]["state"] + == Invitation.State.EXPIRED.value + ) + + @pytest.mark.parametrize( + "filter_name, filter_value, expected_invitations", + [ + ("state", "pending", {"active"}), + ("state", "expired", {"expired", "lapsed"}), + ("state", "accepted", set()), + ("state__in", "pending", {"active"}), + ("state__in", "expired", {"expired", "lapsed"}), + ("state__in", "pending,expired", {"active", "expired", "lapsed"}), + ("state__in", "accepted,revoked", set()), + ], + ) + def test_invitations_filter_state_treats_lapsed_pending_as_expired( + self, + authenticated_client, + create_test_user, + tenants_fixture, + invitations_fixture, + filter_name, + filter_value, + expected_invitations, + ): + active_invitation, expired_invitation = invitations_fixture + lapsed_invitation = self._create_lapsed_invitation( + "lapsed@prowler.com", tenants_fixture[0], create_test_user + ) + invitation_ids = { + "active": str(active_invitation.id), + "expired": str(expired_invitation.id), + "lapsed": str(lapsed_invitation.id), + } + + response = authenticated_client.get( + reverse("invitation-list"), {f"filter[{filter_name}]": filter_value} + ) + + assert response.status_code == status.HTTP_200_OK + assert {invitation["id"] for invitation in response.json()["data"]} == { + invitation_ids[name] for name in expected_invitations + } + @pytest.mark.parametrize( "email", [ @@ -8791,8 +8975,10 @@ class TestInvitationViewSet: "invalid_email@", # There is a pending invitation with this email "testing@prowler.com", + "TESTING@prowler.com", # User is already a member of the tenant TEST_USER, + TEST_USER.upper(), ], ) def test_invitations_create_invalid_email( @@ -9047,6 +9233,56 @@ class TestInvitationViewSet: == "This invitation cannot be revoked." ) + def test_invitations_delete_lapsed_invitation( + self, authenticated_client, invitations_fixture + ): + invitation, *_ = invitations_fixture + invitation.expires_at = datetime.now(UTC) - timedelta(days=1) + 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]["detail"] + == "This invitation cannot be revoked." + ) + invitation.refresh_from_db() + assert invitation.state == Invitation.State.PENDING + + def test_invitations_partial_update_lapsed_invitation( + self, authenticated_client, invitations_fixture + ): + invitation, *_ = invitations_fixture + invitation.expires_at = datetime.now(UTC) - timedelta(days=1) + invitation.save() + data = { + "data": { + "id": str(invitation.id), + "type": "invitations", + "attributes": { + "email": invitation.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]["detail"] + == "This invitation cannot be updated." + ) + invitation.refresh_from_db() + assert invitation.is_lapsed + def test_invitations_accept_invitation_new_user(self, client, invitations_fixture): invitation, *_ = invitations_fixture diff --git a/api/src/backend/api/v1/serializers.py b/api/src/backend/api/v1/serializers.py index 5e1f2fa6d8..271d3a9b47 100644 --- a/api/src/backend/api/v1/serializers.py +++ b/api/src/backend/api/v1/serializers.py @@ -2149,6 +2149,12 @@ class InvitationSerializer(RLSSerializer): if tenant_id is not None: self.fields["roles"].queryset = Role.objects.filter(tenant_id=tenant_id) + def to_representation(self, instance): + data = super().to_representation(instance) + if instance.is_lapsed: + data["state"] = Invitation.State.EXPIRED.value + return data + class Meta: model = Invitation fields = [ @@ -2175,6 +2181,7 @@ class InvitationBaseWriteSerializer(BaseWriteSerializer): self.fields["roles"].queryset = Role.objects.filter(tenant_id=tenant_id) def validate_email(self, value): + value = value.strip().lower() 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(): @@ -2182,9 +2189,13 @@ class InvitationBaseWriteSerializer(BaseWriteSerializer): "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(): + pending_invitations = Invitation.objects.filter( + tenant_id=tenant_id, email=value, state=Invitation.State.PENDING + ) + pending_invitations.filter(Invitation.lapsed_q()).update( + state=Invitation.State.EXPIRED + ) + if pending_invitations.filter(expires_at__gt=datetime.now(UTC)).exists(): raise ValidationError( "Unable to process your request. Please check the information provided and " "try again." diff --git a/api/src/backend/api/v1/views.py b/api/src/backend/api/v1/views.py index 9df5be424c..d0294f4abd 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -4471,7 +4471,7 @@ class InvitationViewSet(BaseRLSViewSet): def partial_update(self, request, *args, **kwargs): instance = self.get_object() - if instance.state != Invitation.State.PENDING: + if instance.state != Invitation.State.PENDING or instance.is_lapsed: raise ValidationError(detail="This invitation cannot be updated.") serializer = self.get_serializer( instance, @@ -4485,7 +4485,7 @@ class InvitationViewSet(BaseRLSViewSet): def destroy(self, request, *args, **kwargs): instance = self.get_object() - if instance.state != Invitation.State.PENDING: + if instance.state != Invitation.State.PENDING or instance.is_lapsed: raise ValidationError(detail="This invitation cannot be revoked.") instance.state = Invitation.State.REVOKED instance.save() diff --git a/ui/changelog.d/invitation-row-actions-non-pending.fixed.md b/ui/changelog.d/invitation-row-actions-non-pending.fixed.md new file mode 100644 index 0000000000..dda0527b0d --- /dev/null +++ b/ui/changelog.d/invitation-row-actions-non-pending.fixed.md @@ -0,0 +1 @@ +Edit and Revoke actions are disabled for expired and revoked invitations diff --git a/ui/components/invitations/table/data-table-row-actions.test.tsx b/ui/components/invitations/table/data-table-row-actions.test.tsx new file mode 100644 index 0000000000..991b6b5974 --- /dev/null +++ b/ui/components/invitations/table/data-table-row-actions.test.tsx @@ -0,0 +1,60 @@ +import { Row } from "@tanstack/react-table"; +import { render, screen } from "@testing-library/react"; +import userEvent from "@testing-library/user-event"; +import { describe, expect, it, vi } from "vitest"; + +vi.mock("next/navigation", () => ({ + useRouter: () => ({ push: vi.fn() }), +})); + +vi.mock("../forms", () => ({ + DeleteForm: () =>
, + EditForm: () => , +})); + +import { InvitationProps } from "@/types"; + +import { DataTableRowActions } from "./data-table-row-actions"; + +const createRow = (state: string) => + ({ + original: { + id: "invitation-1", + attributes: { email: "jane@example.com", state }, + relationships: { inviter: { data: { type: "users", id: "user-1" } } }, + }, + }) as unknown as Row