fix(invitations): expire lapsed invitations on re-invite (#12831)

This commit is contained in:
Pedro Martín
2026-09-17 12:12:14 +02:00
committed by GitHub
parent 396ccf56bb
commit f382400037
9 changed files with 346 additions and 19 deletions
@@ -0,0 +1 @@
Lapsed pending invitations are reported as expired and no longer block a new invitation for the same email
+14 -2
View File
@@ -1440,8 +1440,20 @@ class InvitationFilter(FilterSet):
inserted_at = DateFilter(field_name="inserted_at", lookup_expr="date") inserted_at = DateFilter(field_name="inserted_at", lookup_expr="date")
updated_at = DateFilter(field_name="updated_at", lookup_expr="date") updated_at = DateFilter(field_name="updated_at", lookup_expr="date")
expires_at = DateFilter(field_name="expires_at", lookup_expr="date") expires_at = DateFilter(field_name="expires_at", lookup_expr="date")
state = ChoiceFilter(choices=Invitation.State.choices) state = ChoiceFilter(choices=Invitation.State.choices, method="filter_state")
state__in = ChoiceInFilter(choices=Invitation.State.choices, lookup_expr="in") 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: class Meta:
model = Invitation model = Invitation
+9
View File
@@ -1380,6 +1380,15 @@ class Invitation(RowLevelSecurityProtectedModel):
self.email = self.email.strip().lower() self.email = self.email.strip().lower()
super().save(*args, **kwargs) 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): class Meta(RowLevelSecurityProtectedModel.Meta):
db_table = "invitations" db_table = "invitations"
+236
View File
@@ -8784,6 +8784,190 @@ class TestInvitationViewSet:
user.id 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( @pytest.mark.parametrize(
"email", "email",
[ [
@@ -8791,8 +8975,10 @@ class TestInvitationViewSet:
"invalid_email@", "invalid_email@",
# There is a pending invitation with this email # There is a pending invitation with this email
"testing@prowler.com", "testing@prowler.com",
"TESTING@prowler.com",
# User is already a member of the tenant # User is already a member of the tenant
TEST_USER, TEST_USER,
TEST_USER.upper(),
], ],
) )
def test_invitations_create_invalid_email( def test_invitations_create_invalid_email(
@@ -9047,6 +9233,56 @@ class TestInvitationViewSet:
== "This invitation cannot be revoked." == "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): def test_invitations_accept_invitation_new_user(self, client, invitations_fixture):
invitation, *_ = invitations_fixture invitation, *_ = invitations_fixture
+14 -3
View File
@@ -2150,6 +2150,12 @@ class InvitationSerializer(RLSSerializer):
if tenant_id is not None: if tenant_id is not None:
self.fields["roles"].queryset = Role.objects.filter(tenant_id=tenant_id) 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: class Meta:
model = Invitation model = Invitation
fields = [ fields = [
@@ -2176,6 +2182,7 @@ class InvitationBaseWriteSerializer(BaseWriteSerializer):
self.fields["roles"].queryset = Role.objects.filter(tenant_id=tenant_id) self.fields["roles"].queryset = Role.objects.filter(tenant_id=tenant_id)
def validate_email(self, value): def validate_email(self, value):
value = value.strip().lower()
user = User.objects.filter(email=value).first() user = User.objects.filter(email=value).first()
tenant_id = self.context["tenant_id"] tenant_id = self.context["tenant_id"]
if user and Membership.objects.filter(user=user, tenant=tenant_id).exists(): 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 " "The user may already be a member of the tenant or there was an issue with the "
"email provided." "email provided."
) )
if Invitation.objects.filter( pending_invitations = Invitation.objects.filter(
email=value, state=Invitation.State.PENDING tenant_id=tenant_id, email=value, state=Invitation.State.PENDING
).exists(): )
pending_invitations.filter(Invitation.lapsed_q()).update(
state=Invitation.State.EXPIRED
)
if pending_invitations.filter(expires_at__gt=datetime.now(UTC)).exists():
raise ValidationError( raise ValidationError(
"Unable to process your request. Please check the information provided and " "Unable to process your request. Please check the information provided and "
"try again." "try again."
+2 -2
View File
@@ -4476,7 +4476,7 @@ class InvitationViewSet(BaseRLSViewSet):
def partial_update(self, request, *args, **kwargs): def partial_update(self, request, *args, **kwargs):
instance = self.get_object() 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.") raise ValidationError(detail="This invitation cannot be updated.")
serializer = self.get_serializer( serializer = self.get_serializer(
instance, instance,
@@ -4490,7 +4490,7 @@ class InvitationViewSet(BaseRLSViewSet):
def destroy(self, request, *args, **kwargs): def destroy(self, request, *args, **kwargs):
instance = self.get_object() 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.") raise ValidationError(detail="This invitation cannot be revoked.")
instance.state = Invitation.State.REVOKED instance.state = Invitation.State.REVOKED
instance.save() instance.save()
@@ -0,0 +1 @@
Edit and Revoke actions are disabled for expired and revoked invitations
@@ -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: () => <div data-testid="delete-form" />,
EditForm: () => <div data-testid="edit-form" />,
}));
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<InvitationProps>;
const openMenu = async (user: ReturnType<typeof userEvent.setup>) => {
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(<DataTableRowActions row={createRow("pending")} />);
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(<DataTableRowActions row={createRow(state)} />);
await openMenu(user);
expect(getMenuItem("Edit Invitation")).toHaveAttribute("data-disabled");
expect(getMenuItem("Revoke Invitation")).toHaveAttribute("data-disabled");
},
);
});
@@ -11,26 +11,23 @@ import {
ActionDropdownItem, ActionDropdownItem,
} from "@/components/shadcn/dropdown"; } from "@/components/shadcn/dropdown";
import { Modal } from "@/components/shadcn/modal"; import { Modal } from "@/components/shadcn/modal";
import { InvitationProps } from "@/types";
import { DeleteForm, EditForm } from "../forms"; import { DeleteForm, EditForm } from "../forms";
interface DataTableRowActionsProps<InvitationProps> { interface DataTableRowActionsProps {
row: Row<InvitationProps>; row: Row<InvitationProps>;
roles?: { id: string; name: string }[]; roles?: { id: string; name: string }[];
} }
export function DataTableRowActions<InvitationProps>({ export function DataTableRowActions({ row, roles }: DataTableRowActionsProps) {
row,
roles,
}: DataTableRowActionsProps<InvitationProps>) {
const router = useRouter(); const router = useRouter();
const [isEditOpen, setIsEditOpen] = useState(false); const [isEditOpen, setIsEditOpen] = useState(false);
const [isDeleteOpen, setIsDeleteOpen] = useState(false); const [isDeleteOpen, setIsDeleteOpen] = useState(false);
const invitationId = (row.original as { id: string }).id; const invitationId = row.original.id;
const invitationEmail = (row.original as any).attributes?.email; const invitationEmail = row.original.attributes.email;
const invitationRole = (row.original as any).relationships?.role?.attributes const invitationRole = row.original.relationships.role?.attributes?.name;
?.name; const isInvitationPending = row.original.attributes.state === "pending";
const invitationAccepted = (row.original as any).attributes?.state;
return ( return (
<> <>
@@ -69,7 +66,7 @@ export function DataTableRowActions<InvitationProps>({
icon={<Pencil />} icon={<Pencil />}
label="Edit Invitation" label="Edit Invitation"
onSelect={() => setIsEditOpen(true)} onSelect={() => setIsEditOpen(true)}
disabled={invitationAccepted === "accepted"} disabled={!isInvitationPending}
/> />
<ActionDropdownDangerZone> <ActionDropdownDangerZone>
<ActionDropdownItem <ActionDropdownItem
@@ -77,7 +74,7 @@ export function DataTableRowActions<InvitationProps>({
label="Revoke Invitation" label="Revoke Invitation"
destructive destructive
onSelect={() => setIsDeleteOpen(true)} onSelect={() => setIsDeleteOpen(true)}
disabled={invitationAccepted === "accepted"} disabled={!isInvitationPending}
/> />
</ActionDropdownDangerZone> </ActionDropdownDangerZone>
</ActionDropdown> </ActionDropdown>