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 8a9f37c366..2b3a4b333a 100644 --- a/api/src/backend/api/filters.py +++ b/api/src/backend/api/filters.py @@ -1440,8 +1440,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 4f1e963c3d..ef19712f94 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 663f199339..639ddade61 100644 --- a/api/src/backend/api/v1/serializers.py +++ b/api/src/backend/api/v1/serializers.py @@ -2150,6 +2150,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 = [ @@ -2176,6 +2182,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(): @@ -2183,9 +2190,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 df89a69523..8839fcc686 100644 --- a/api/src/backend/api/v1/views.py +++ b/api/src/backend/api/v1/views.py @@ -4476,7 +4476,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, @@ -4490,7 +4490,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; + +const openMenu = async (user: ReturnType) => { + await user.click(screen.getByRole("button", { name: "Open actions menu" })); +}; + +const getMenuItem = (label: string) => + screen.getByText(label).closest("[role='menuitem']"); + +describe("DataTableRowActions (invitations)", () => { + it("enables Edit and Revoke for a pending invitation", async () => { + const user = userEvent.setup(); + render(); + + await openMenu(user); + + expect(getMenuItem("Edit Invitation")).not.toHaveAttribute("data-disabled"); + expect(getMenuItem("Revoke Invitation")).not.toHaveAttribute( + "data-disabled", + ); + }); + + it.each(["accepted", "expired", "revoked"])( + "disables Edit and Revoke for a %s invitation", + async (state) => { + const user = userEvent.setup(); + render(); + + await openMenu(user); + + expect(getMenuItem("Edit Invitation")).toHaveAttribute("data-disabled"); + expect(getMenuItem("Revoke Invitation")).toHaveAttribute("data-disabled"); + }, + ); +}); diff --git a/ui/components/invitations/table/data-table-row-actions.tsx b/ui/components/invitations/table/data-table-row-actions.tsx index a0ef693279..b2085fc9b3 100644 --- a/ui/components/invitations/table/data-table-row-actions.tsx +++ b/ui/components/invitations/table/data-table-row-actions.tsx @@ -11,26 +11,23 @@ import { ActionDropdownItem, } from "@/components/shadcn/dropdown"; import { Modal } from "@/components/shadcn/modal"; +import { InvitationProps } from "@/types"; import { DeleteForm, EditForm } from "../forms"; -interface DataTableRowActionsProps { +interface DataTableRowActionsProps { row: Row; roles?: { id: string; name: string }[]; } -export function DataTableRowActions({ - row, - roles, -}: DataTableRowActionsProps) { +export function DataTableRowActions({ row, roles }: DataTableRowActionsProps) { const router = useRouter(); const [isEditOpen, setIsEditOpen] = useState(false); const [isDeleteOpen, setIsDeleteOpen] = useState(false); - const invitationId = (row.original as { id: string }).id; - const invitationEmail = (row.original as any).attributes?.email; - const invitationRole = (row.original as any).relationships?.role?.attributes - ?.name; - const invitationAccepted = (row.original as any).attributes?.state; + const invitationId = row.original.id; + const invitationEmail = row.original.attributes.email; + const invitationRole = row.original.relationships.role?.attributes?.name; + const isInvitationPending = row.original.attributes.state === "pending"; return ( <> @@ -69,7 +66,7 @@ export function DataTableRowActions({ icon={} label="Edit Invitation" onSelect={() => setIsEditOpen(true)} - disabled={invitationAccepted === "accepted"} + disabled={!isInvitationPending} /> ({ label="Revoke Invitation" destructive onSelect={() => setIsDeleteOpen(true)} - disabled={invitationAccepted === "accepted"} + disabled={!isInvitationPending} />