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

Co-authored-by: Pedro Martín <pedromarting3@gmail.com>
This commit is contained in:
Prowler Bot
2026-09-17 12:29:09 +02:00
committed by GitHub
co-authored by Pedro Martín
parent 8ae8ffab4a
commit eabf10b109
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
@@ -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
+9
View File
@@ -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"
+236
View File
@@ -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
+14 -3
View File
@@ -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."
+2 -2
View File
@@ -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()