From 076947913b8dc44d474bef0aea0fc358a0f193cc Mon Sep 17 00:00:00 2001 From: Artur Shiriev Date: Sat, 10 Oct 2026 13:00:52 +0300 Subject: [PATCH] fix: revoke the JWT on logout via a jti denylist --- app/api/auth.py | 16 ++++++++ app/api/endpoints/auth.py | 11 +++++- app/database/tables.py | 9 ++++- app/ioc.py | 9 +++++ app/repositories/revoked_tokens_repository.py | 28 +++++++++++++ app/use_cases/check_token_revoked.py | 14 +++++++ app/use_cases/revoke_token.py | 19 +++++++++ docs/adr/0011-logout-revokes-by-jti.md | 14 +++++++ .../versions/2026-10-10_revoked_tokens.py | 36 +++++++++++++++++ tests/api/test_auth_api.py | 39 ++++++++++++++++--- tests/use_cases/test_revoke_token.py | 25 ++++++++++++ 11 files changed, 212 insertions(+), 8 deletions(-) create mode 100644 app/repositories/revoked_tokens_repository.py create mode 100644 app/use_cases/check_token_revoked.py create mode 100644 app/use_cases/revoke_token.py create mode 100644 docs/adr/0011-logout-revokes-by-jti.md create mode 100644 migrations/versions/2026-10-10_revoked_tokens.py create mode 100644 tests/use_cases/test_revoke_token.py diff --git a/app/api/auth.py b/app/api/auth.py index 74c3921..a00f18f 100644 --- a/app/api/auth.py +++ b/app/api/auth.py @@ -1,12 +1,16 @@ import datetime as dt import typing +import uuid import litestar +import modern_di from litestar.config.app import AppConfig from litestar.connection import ASGIConnection from litestar.plugins import InitPlugin from litestar.security.jwt import JWTCookieAuth, Token +from modern_di_litestar import fetch_di_container +from app import ioc from app.actor import Actor from app.settings import settings @@ -21,8 +25,20 @@ async def retrieve_user_handler(token: Token, _connection: ASGIConnection) -> Ac return None +async def revoked_token_handler(token: Token, connection: ASGIConnection) -> bool: + async with fetch_di_container(connection.app).build_child_container(scope=modern_di.Scope.REQUEST) as container: + check_token_revoked: typing.Final = container.resolve_provider(ioc.UseCases.check_token_revoked_use_case) + return await check_token_revoked(jti=str(token.jti)) + + +def new_token_id() -> str: + return str(uuid.uuid4()) + + jwt_cookie_auth: typing.Final = JWTCookieAuth[Actor]( retrieve_user_handler=retrieve_user_handler, + revoked_token_handler=revoked_token_handler, + require_claims=["jti"], token_secret=settings.jwt_secret, default_token_expiration=dt.timedelta(seconds=settings.jwt_lifetime_seconds), secure=settings.jwt_cookie_secure, diff --git a/app/api/endpoints/auth.py b/app/api/endpoints/auth.py index 2131a9a..fab2403 100644 --- a/app/api/endpoints/auth.py +++ b/app/api/endpoints/auth.py @@ -6,11 +6,12 @@ from litestar.exceptions import NotAuthorizedException from litestar.response import Response -from app.api.auth import AuthedRequest, jwt_cookie_auth +from app.api.auth import AuthedRequest, jwt_cookie_auth, new_token_id from app.schemas import api as schemas from app.use_cases.authenticate_user import AuthenticateUserUseCase from app.use_cases.fetch_user import FetchUserUseCase from app.use_cases.register_user import RegisterUserUseCase +from app.use_cases.revoke_token import RevokeTokenUseCase @litestar.post("/auth/register/", status_code=status_codes.HTTP_201_CREATED, exclude_from_auth=True) @@ -21,6 +22,7 @@ async def register( user: typing.Final = await register_user_use_case(data=data) return jwt_cookie_auth.login( identifier=str(user.id), + token_unique_jwt_id=new_token_id(), response_body=schemas.User.model_validate(user), response_status_code=status_codes.HTTP_201_CREATED, ) @@ -36,13 +38,18 @@ async def login( raise NotAuthorizedException(detail="Invalid username or password") return jwt_cookie_auth.login( identifier=str(user.id), + token_unique_jwt_id=new_token_id(), response_body=schemas.User.model_validate(user), response_status_code=status_codes.HTTP_200_OK, ) @litestar.post("/auth/logout/", status_code=status_codes.HTTP_204_NO_CONTENT) -async def logout() -> Response[None]: +async def logout( + request: AuthedRequest, + revoke_token_use_case: NamedDependency[RevokeTokenUseCase], +) -> Response[None]: + await revoke_token_use_case(jti=str(request.auth.jti), expires_at=request.auth.exp) response: typing.Final = Response(content=None, status_code=status_codes.HTTP_204_NO_CONTENT) response.delete_cookie(jwt_cookie_auth.key) return response diff --git a/app/database/tables.py b/app/database/tables.py index be7df5e..2538446 100644 --- a/app/database/tables.py +++ b/app/database/tables.py @@ -4,7 +4,7 @@ import uuid import sqlalchemy as sa -from advanced_alchemy.base import BigIntAuditBase, BigIntBase, orm_registry +from advanced_alchemy.base import BigIntAuditBase, BigIntBase, DefaultBase, orm_registry from advanced_alchemy.types import GUID, DateTimeUTC from sqlalchemy import orm @@ -94,3 +94,10 @@ class MessagesTable(BigIntBase): ) edited_at: orm.Mapped[dt.datetime | None] = orm.mapped_column(DateTimeUTC(timezone=True), nullable=True) deleted_at: orm.Mapped[dt.datetime | None] = orm.mapped_column(DateTimeUTC(timezone=True), nullable=True) + + +class RevokedTokensTable(DefaultBase): + __tablename__ = "revoked_tokens" + + jti: orm.Mapped[str] = orm.mapped_column(sa.String(length=36), primary_key=True) + expires_at: orm.Mapped[dt.datetime] = orm.mapped_column(DateTimeUTC(timezone=True), index=True) diff --git a/app/ioc.py b/app/ioc.py index 41df3a1..eea6a5f 100644 --- a/app/ioc.py +++ b/app/ioc.py @@ -7,8 +7,10 @@ from app.repositories.chat_members_repository import ChatMembersRepository from app.repositories.chats_repository import ChatsRepository from app.repositories.messages_repository import MessagesRepository +from app.repositories.revoked_tokens_repository import RevokedTokensRepository from app.repositories.users_repository import UsersRepository from app.use_cases.authenticate_user import AuthenticateUserUseCase +from app.use_cases.check_token_revoked import CheckTokenRevokedUseCase from app.use_cases.create_chat import CreateChatUseCase from app.use_cases.create_message import CreateMessageUseCase from app.use_cases.delete_message import DeleteMessageUseCase @@ -19,6 +21,7 @@ from app.use_cases.fetch_user import FetchUserUseCase from app.use_cases.mark_read import MarkReadUseCase from app.use_cases.register_user import RegisterUserUseCase +from app.use_cases.revoke_token import RevokeTokenUseCase class Database(Group): @@ -55,6 +58,10 @@ class Repositories(Group, scope=Scope.REQUEST): creator=MessagesRepository, kwargs={"session": Database.database_session, "auto_commit": False}, ) + revoked_tokens_repository = providers.Factory( + creator=RevokedTokensRepository, + kwargs={"session": Database.database_session, "auto_commit": False}, + ) class UseCases(Group, scope=Scope.REQUEST): @@ -69,6 +76,8 @@ class UseCases(Group, scope=Scope.REQUEST): fetch_chats_use_case = providers.Factory(creator=FetchChatsUseCase) fetch_user_use_case = providers.Factory(creator=FetchUserUseCase) mark_read_use_case = providers.Factory(creator=MarkReadUseCase) + revoke_token_use_case = providers.Factory(creator=RevokeTokenUseCase) + check_token_revoked_use_case = providers.Factory(creator=CheckTokenRevokedUseCase) ALL_GROUPS: typing.Final[list[type[Group]]] = [Database, Repositories, UseCases] diff --git a/app/repositories/revoked_tokens_repository.py b/app/repositories/revoked_tokens_repository.py new file mode 100644 index 0000000..f226fc1 --- /dev/null +++ b/app/repositories/revoked_tokens_repository.py @@ -0,0 +1,28 @@ +import datetime as dt + +from advanced_alchemy.repository import SQLAlchemyAsyncRepository +from advanced_alchemy.service import SQLAlchemyAsyncRepositoryService +from sqlalchemy.dialects import postgresql + +from app.database import tables + + +class RevokedTokensRepository(SQLAlchemyAsyncRepositoryService[tables.RevokedTokensTable]): + class BaseRepository(SQLAlchemyAsyncRepository[tables.RevokedTokensTable]): + model_type = tables.RevokedTokensTable + id_attribute = "jti" + + repository_type = BaseRepository + + async def is_revoked(self, jti: str) -> bool: + return await self.exists(jti=jti) + + async def revoke(self, jti: str, expires_at: dt.datetime) -> None: + await self.repository.session.execute( + postgresql.insert(tables.RevokedTokensTable) + .values(jti=jti, expires_at=expires_at) + .on_conflict_do_nothing(index_elements=[tables.RevokedTokensTable.jti]) + ) + + async def prune_expired(self, now: dt.datetime) -> None: + await self.delete_where(tables.RevokedTokensTable.expires_at < now) diff --git a/app/use_cases/check_token_revoked.py b/app/use_cases/check_token_revoked.py new file mode 100644 index 0000000..396bfbf --- /dev/null +++ b/app/use_cases/check_token_revoked.py @@ -0,0 +1,14 @@ +import dataclasses + +from db_retry import postgres_retry + +from app.repositories.revoked_tokens_repository import RevokedTokensRepository + + +@dataclasses.dataclass(kw_only=True, frozen=True, slots=True) +class CheckTokenRevokedUseCase: + revoked_tokens_repository: RevokedTokensRepository + + @postgres_retry + async def __call__(self, *, jti: str) -> bool: + return await self.revoked_tokens_repository.is_revoked(jti) diff --git a/app/use_cases/revoke_token.py b/app/use_cases/revoke_token.py new file mode 100644 index 0000000..96d4536 --- /dev/null +++ b/app/use_cases/revoke_token.py @@ -0,0 +1,19 @@ +import dataclasses +import datetime as dt + +from db_retry import Transaction, postgres_retry + +from app.repositories.revoked_tokens_repository import RevokedTokensRepository + + +@dataclasses.dataclass(kw_only=True, frozen=True, slots=True) +class RevokeTokenUseCase: + transaction: Transaction + revoked_tokens_repository: RevokedTokensRepository + + @postgres_retry + async def __call__(self, *, jti: str, expires_at: dt.datetime) -> None: + async with self.transaction: + await self.revoked_tokens_repository.prune_expired(dt.datetime.now(tz=dt.UTC)) + await self.revoked_tokens_repository.revoke(jti, expires_at) + await self.transaction.commit() diff --git a/docs/adr/0011-logout-revokes-by-jti.md b/docs/adr/0011-logout-revokes-by-jti.md new file mode 100644 index 0000000..6283ea1 --- /dev/null +++ b/docs/adr/0011-logout-revokes-by-jti.md @@ -0,0 +1,14 @@ +# Logout revokes the token by its jti + +Logout writes the token's `jti` and expiry to `revoked_tokens`, and +`revoked_token_handler` rejects any token listed there. A copy taken before +logout stops working, and the user's other sessions keep theirs. A token +version on `users` was rejected because one logout would end every session of +that user, and short-lived tokens with refresh because they replace the whole +auth model. The cost is one primary-key lookup per authenticated request, on a +session the middleware closes before the handler opens its own, so a request +still holds one pooled connection at a time. That lookup is the read +[ADR-0010](0010-auth-carries-an-actor-id.md) kept out of authentication; the +user row is still never loaded. A token without a `jti`, which includes every +token issued before this change, gets a 401. Each logout prunes expired rows, +so the table holds at most a token lifetime of logouts. diff --git a/migrations/versions/2026-10-10_revoked_tokens.py b/migrations/versions/2026-10-10_revoked_tokens.py new file mode 100644 index 0000000..cdc00c0 --- /dev/null +++ b/migrations/versions/2026-10-10_revoked_tokens.py @@ -0,0 +1,36 @@ +"""revoked tokens. + +Revision ID: 8cf557a6bd3a +Revises: fa15d87677c3 +Create Date: 2026-10-10 09:58:04.832132 + +""" + +import advanced_alchemy +import sqlalchemy as sa +from alembic import op + + +revision = "8cf557a6bd3a" +down_revision = "fa15d87677c3" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.create_table( + "revoked_tokens", + sa.Column("jti", sa.String(length=36), nullable=False), + sa.Column("expires_at", advanced_alchemy.types.datetime.DateTimeUTC(timezone=True), nullable=False), + sa.PrimaryKeyConstraint("jti", name=op.f("pk_revoked_tokens")), + ) + op.create_index(op.f("ix_revoked_tokens_expires_at"), "revoked_tokens", ["expires_at"], unique=False) + # ### end Alembic commands ### + + +def downgrade() -> None: + # ### commands auto generated by Alembic - please adjust! ### + op.drop_index(op.f("ix_revoked_tokens_expires_at"), table_name="revoked_tokens") + op.drop_table("revoked_tokens") + # ### end Alembic commands ### diff --git a/tests/api/test_auth_api.py b/tests/api/test_auth_api.py index f605ada..478496d 100644 --- a/tests/api/test_auth_api.py +++ b/tests/api/test_auth_api.py @@ -7,7 +7,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from app.actor import Actor -from app.api.auth import jwt_cookie_auth, retrieve_user_handler +from app.api.auth import jwt_cookie_auth, new_token_id, retrieve_user_handler from app.database import tables @@ -83,6 +83,36 @@ async def test_logout_clears_the_cookie(client: AsyncClient) -> None: assert me_response.status_code == 401 +@pytest.mark.usefixtures("db_session") +async def test_logout_revokes_a_copy_of_the_token(client: AsyncClient) -> None: + await client.post("/api/auth/register/", json=REGISTRATION) + copied_token = client.cookies[jwt_cookie_auth.key] + await client.post("/api/auth/logout/") + client.cookies.set(jwt_cookie_auth.key, copied_token) + response = await client.get("/api/auth/me/") + assert response.status_code == 401 + + +@pytest.mark.usefixtures("db_session") +async def test_logout_keeps_other_sessions(client: AsyncClient) -> None: + await client.post("/api/auth/register/", json=REGISTRATION) + other_session_token = client.cookies[jwt_cookie_auth.key] + await client.post("/api/auth/login/", json={"username": "alice", "password": "hunter2hunter2"}) + await client.post("/api/auth/logout/") + client.cookies.set(jwt_cookie_auth.key, other_session_token) + response = await client.get("/api/auth/me/") + assert response.status_code == 200 + + +@pytest.mark.usefixtures("db_session") +async def test_me_rejects_token_without_jti(client: AsyncClient) -> None: + response = await client.post("/api/auth/register/", json=REGISTRATION) + token = jwt_cookie_auth.create_token(identifier=str(response.json()["id"])) + client.cookies.set(jwt_cookie_auth.key, token) + response = await client.get("/api/auth/me/") + assert response.status_code == 401 + + async def test_password_is_stored_hashed(client: AsyncClient, db_session: AsyncSession) -> None: await client.post("/api/auth/register/", json=REGISTRATION) stored = await db_session.scalar(sa.select(tables.UsersTable.password_hash)) @@ -104,7 +134,7 @@ async def test_me_rejects_token_with_non_numeric_subject(client: AsyncClient) -> inside auth middleware, so an uncaught ValueError there is an unhandled server error on a request an attacker fully controls the token for. """ - token = jwt_cookie_auth.create_token(identifier="not-a-number") + token = jwt_cookie_auth.create_token(identifier="not-a-number", token_unique_jwt_id=new_token_id()) client.cookies.set(jwt_cookie_auth.key, token) response = await client.get("/api/auth/me/") assert response.status_code == 401 @@ -122,7 +152,7 @@ async def test_metrics_are_reachable_without_a_cookie(client: AsyncClient) -> No @pytest.mark.usefixtures("db_session") async def test_me_rejects_token_for_a_user_that_no_longer_exists(client: AsyncClient) -> None: - token = jwt_cookie_auth.create_token(identifier="999999999") + token = jwt_cookie_auth.create_token(identifier="999999999", token_unique_jwt_id=new_token_id()) client.cookies.set(jwt_cookie_auth.key, token) response = await client.get("/api/auth/me/") assert response.status_code == 401 @@ -138,8 +168,7 @@ async def test_retrieve_user_handler_resolves_an_actor_without_reading_the_datab The cost this accepts: authentication no longer proves the user row exists. A token that outlives its user still authenticates - reads come back empty, writes hit the messages.user_id - foreign key. Nothing can reach that state today; a delete-user path would have to solve it - alongside logout not revoking the JWT. + foreign key. Nothing can reach that state today; a delete-user path would have to solve it. """ token = Token(sub="42", exp=dt.datetime.now(tz=dt.UTC) + dt.timedelta(minutes=5)) assert await retrieve_user_handler(token, None) == Actor(id=42) # ty: ignore[invalid-argument-type] diff --git a/tests/use_cases/test_revoke_token.py b/tests/use_cases/test_revoke_token.py new file mode 100644 index 0000000..2de28dc --- /dev/null +++ b/tests/use_cases/test_revoke_token.py @@ -0,0 +1,25 @@ +import datetime as dt + +import sqlalchemy as sa +from sqlalchemy.ext.asyncio import AsyncSession + +from app.database import tables +from app.use_cases.check_token_revoked import CheckTokenRevokedUseCase +from app.use_cases.revoke_token import RevokeTokenUseCase + + +async def test_revoke_prunes_expired_rows(revoke_token_use_case: RevokeTokenUseCase, db_session: AsyncSession) -> None: + now = dt.datetime.now(tz=dt.UTC) + await revoke_token_use_case(jti="expired", expires_at=now - dt.timedelta(seconds=1)) + await revoke_token_use_case(jti="live", expires_at=now + dt.timedelta(hours=1)) + stored = await db_session.scalars(sa.select(tables.RevokedTokensTable.jti)) + assert stored.all() == ["live"] + + +async def test_revoking_twice_is_a_no_op( + revoke_token_use_case: RevokeTokenUseCase, check_token_revoked_use_case: CheckTokenRevokedUseCase +) -> None: + expires_at = dt.datetime.now(tz=dt.UTC) + dt.timedelta(hours=1) + await revoke_token_use_case(jti="twice", expires_at=expires_at) + await revoke_token_use_case(jti="twice", expires_at=expires_at) + assert await check_token_revoked_use_case(jti="twice")