From 18b3c492343f66c1773c234dc346ac7a3c4551cc Mon Sep 17 00:00:00 2001 From: =?UTF-8?q?Adri=C3=A1n=20Pe=C3=B1a?= Date: Thu, 9 Jul 2026 10:36:58 +0200 Subject: [PATCH] fix(api): invalidate tokens after password updates (#11901) --- .../password-token-refresh.fixed.md | 1 + api/src/backend/api/sse/channelmanager.py | 3 +- .../tests/integration/test_authentication.py | 119 ++++++++++++++++++ .../api/v1/serializer_utils/authentication.py | 18 +++ api/src/backend/api/v1/serializers.py | 25 +++- api/src/backend/config/django/base.py | 1 + 6 files changed, 164 insertions(+), 3 deletions(-) create mode 100644 api/changelog.d/password-token-refresh.fixed.md create mode 100644 api/src/backend/api/v1/serializer_utils/authentication.py diff --git a/api/changelog.d/password-token-refresh.fixed.md b/api/changelog.d/password-token-refresh.fixed.md new file mode 100644 index 0000000000..87101d2d52 --- /dev/null +++ b/api/changelog.d/password-token-refresh.fixed.md @@ -0,0 +1 @@ +Session tokens are rejected after account password updates diff --git a/api/src/backend/api/sse/channelmanager.py b/api/src/backend/api/sse/channelmanager.py index 9190d4ab16..84a9362a19 100644 --- a/api/src/backend/api/sse/channelmanager.py +++ b/api/src/backend/api/sse/channelmanager.py @@ -16,7 +16,7 @@ if TYPE_CHECKING: class SSEChannelManager(DefaultChannelManager): """Connect `django-eventstream` to the platform's SSE viewsets.""" - def get_channels_for_request(self, request: Request, view_kwargs: dict) -> set[str]: # noqa: vulture + def get_channels_for_request(self, request: Request, view_kwargs: dict) -> set[str]: """Return the request's channels scoped to the active JWT tenant. Args: @@ -30,6 +30,7 @@ class SSEChannelManager(DefaultChannelManager): The subset of `request.sse_channels` whose embedded tenant matches the active request tenant. """ + _ = view_kwargs try: request_tenant_id = UUID(str(getattr(request, "tenant_id", None))) except (TypeError, ValueError): diff --git a/api/src/backend/api/tests/integration/test_authentication.py b/api/src/backend/api/tests/integration/test_authentication.py index c68d95d2b6..23cf07afcb 100644 --- a/api/src/backend/api/tests/integration/test_authentication.py +++ b/api/src/backend/api/tests/integration/test_authentication.py @@ -1,3 +1,4 @@ +import json import time from datetime import UTC, datetime, timedelta from uuid import uuid4 @@ -8,6 +9,10 @@ from conftest import TEST_PASSWORD, get_api_tokens, get_authorization_header from django.urls import reverse from drf_simple_apikey.crypto import get_crypto from rest_framework.test import APIClient +from rest_framework_simplejwt.token_blacklist.models import ( + BlacklistedToken, + OutstandingToken, +) @pytest.mark.django_db @@ -103,6 +108,118 @@ def test_refresh_token(create_test_user, tenants_fixture): assert new_refresh_response.status_code == 200 +@pytest.mark.django_db +def test_password_change_invalidates_existing_tokens(create_test_user, tenants_fixture): + client = APIClient() + new_password = "ChangedSecret123@" + + access_token, refresh_token = get_api_tokens( + client, create_test_user.email, TEST_PASSWORD + ) + auth_headers = get_authorization_header(access_token) + outstanding_token_ids = list( + OutstandingToken.objects.filter(user=create_test_user).values_list( + "id", flat=True + ) + ) + assert outstanding_token_ids + assert not BlacklistedToken.objects.filter( + token_id__in=outstanding_token_ids + ).exists() + + password_change_payload = { + "data": { + "type": "users", + "id": str(create_test_user.id), + "attributes": {"password": new_password}, + } + } + password_change_response = client.patch( + reverse("user-detail", kwargs={"pk": create_test_user.id}), + data=json.dumps(password_change_payload), + headers=auth_headers, + content_type="application/vnd.api+json", + ) + assert password_change_response.status_code == 200, password_change_response.json() + assert BlacklistedToken.objects.filter( + token_id__in=outstanding_token_ids + ).count() == len(outstanding_token_ids) + + old_access_response = client.get(reverse("user-me"), headers=auth_headers) + assert old_access_response.status_code == 401 + + old_refresh_response = client.post( + reverse("token-refresh"), + data={ + "data": { + "type": "tokens-refresh", + "attributes": {"refresh": refresh_token}, + } + }, + format="vnd.api+json", + ) + assert old_refresh_response.status_code == 400 + + new_access_token, _ = get_api_tokens(client, create_test_user.email, new_password) + new_access_response = client.get( + reverse("user-me"), headers=get_authorization_header(new_access_token) + ) + assert new_access_response.status_code == 200 + + +@pytest.mark.django_db +def test_password_change_invalidates_rotated_refresh_token( + create_test_user, tenants_fixture +): + client = APIClient() + new_password = "ChangedSecret123@" + + access_token, refresh_token = get_api_tokens( + client, create_test_user.email, TEST_PASSWORD + ) + rotated_refresh_response = client.post( + reverse("token-refresh"), + data={ + "data": { + "type": "tokens-refresh", + "attributes": {"refresh": refresh_token}, + } + }, + format="vnd.api+json", + ) + assert rotated_refresh_response.status_code == 200 + rotated_refresh_token = rotated_refresh_response.json()["data"]["attributes"][ + "refresh" + ] + + password_change_payload = { + "data": { + "type": "users", + "id": str(create_test_user.id), + "attributes": {"password": new_password}, + } + } + password_change_response = client.patch( + reverse("user-detail", kwargs={"pk": create_test_user.id}), + data=json.dumps(password_change_payload), + headers=get_authorization_header(access_token), + content_type="application/vnd.api+json", + ) + assert password_change_response.status_code == 200, password_change_response.json() + + old_rotated_refresh_response = client.post( + reverse("token-refresh"), + data={ + "data": { + "type": "tokens-refresh", + "attributes": {"refresh": rotated_refresh_token}, + } + }, + format="vnd.api+json", + ) + assert old_rotated_refresh_response.status_code == 400 + + @pytest.mark.django_db def test_user_me_when_inviting_users(create_test_user, tenants_fixture, roles_fixture): client = APIClient() @@ -189,6 +306,7 @@ def test_user_me_when_inviting_users(create_test_user, tenants_fixture, roles_fi class TestTokenSwitchTenant: def test_switch_tenant_with_valid_token(self, tenants_fixture, aws_provider): client = APIClient() + assert aws_provider test_user = "test_email@prowler.com" test_password = "Test_password1@" @@ -1403,6 +1521,7 @@ class TestAPIKeyMultiTenantWorkflows: Verifies RLS enforcement after authentication ensures tenant isolation. """ client = APIClient() + assert aws_provider user1 = User.objects.create_user( name="tenant1_user", diff --git a/api/src/backend/api/v1/serializer_utils/authentication.py b/api/src/backend/api/v1/serializer_utils/authentication.py new file mode 100644 index 0000000000..840cb44a0a --- /dev/null +++ b/api/src/backend/api/v1/serializer_utils/authentication.py @@ -0,0 +1,18 @@ +from api.db_router import MainRouter +from rest_framework_simplejwt.token_blacklist.models import ( + BlacklistedToken, + OutstandingToken, +) + + +def blacklist_user_refresh_tokens(user_id): + outstanding_token_ids = list( + OutstandingToken.objects.using(MainRouter.admin_db) + .filter(user_id=user_id) + .values_list("id", flat=True) + ) + if outstanding_token_ids: + BlacklistedToken.objects.using(MainRouter.admin_db).bulk_create( + [BlacklistedToken(token_id=token_id) for token_id in outstanding_token_ids], + ignore_conflicts=True, + ) diff --git a/api/src/backend/api/v1/serializers.py b/api/src/backend/api/v1/serializers.py index 57aa32f0a0..f650a1368a 100644 --- a/api/src/backend/api/v1/serializers.py +++ b/api/src/backend/api/v1/serializers.py @@ -38,6 +38,7 @@ from api.models import ( UserRoleRelationship, ) from api.rls import Tenant +from api.v1.serializer_utils.authentication import blacklist_user_refresh_tokens from api.v1.serializer_utils.integrations import ( AWSCredentialSerializer, IntegrationConfigField, @@ -61,7 +62,7 @@ from django.contrib.auth import authenticate from django.contrib.auth.models import update_last_login from django.contrib.auth.password_validation import validate_password from django.core.exceptions import ValidationError as DjangoValidationError -from django.db import IntegrityError +from django.db import IntegrityError, transaction from drf_spectacular.utils import extend_schema_field from jwt.exceptions import InvalidKeyError from prowler.lib.mutelist.mutelist import Mutelist @@ -72,7 +73,9 @@ from rest_framework_json_api.relations import SerializerMethodResourceRelatedFie from rest_framework_json_api.serializers import ValidationError from rest_framework_simplejwt.exceptions import TokenError from rest_framework_simplejwt.serializers import TokenObtainPairSerializer +from rest_framework_simplejwt.settings import api_settings from rest_framework_simplejwt.tokens import RefreshToken +from rest_framework_simplejwt.utils import get_md5_hash_password # Base @@ -232,6 +235,18 @@ class TokenRefreshSerializer(BaseSerializerV1): try: # Validate the refresh token refresh = RefreshToken(refresh_token) + if api_settings.CHECK_REVOKE_TOKEN: + user_id = refresh.payload.get(api_settings.USER_ID_CLAIM) + try: + user = User.objects.using(MainRouter.admin_db).get( + **{api_settings.USER_ID_FIELD: user_id} + ) + except User.DoesNotExist: + raise TokenError("User not found.") from None + if refresh.get(api_settings.REVOKE_TOKEN_CLAIM) != ( + get_md5_hash_password(user.password) + ): + raise TokenError("The user's password has been changed.") # Generate new access token access_token = refresh.access_token @@ -405,7 +420,13 @@ class UserUpdateSerializer(BaseWriteSerializer): password = validated_data.pop("password", None) if password: validate_password(password, user=instance) - instance.set_password(password) + with transaction.atomic(using=MainRouter.admin_db): + instance.set_password(password) + for attr, value in validated_data.items(): + setattr(instance, attr, value) + blacklist_user_refresh_tokens(instance.id) + instance.save(using=MainRouter.admin_db) + return instance return super().update(instance, validated_data) diff --git a/api/src/backend/config/django/base.py b/api/src/backend/config/django/base.py index 75d5e6112e..63dab22c2c 100644 --- a/api/src/backend/config/django/base.py +++ b/api/src/backend/config/django/base.py @@ -230,6 +230,7 @@ SIMPLE_JWT = { "JTI_CLAIM": "jti", "USER_ID_FIELD": "id", "USER_ID_CLAIM": "sub", + "CHECK_REVOKE_TOKEN": True, # Issuer and Audience claims, for the moment we will keep these values as default values, they may change in the # future. "AUDIENCE": env.str("DJANGO_JWT_AUDIENCE", "https://api.prowler.com"),