Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 16 additions & 0 deletions app/api/auth.py
Original file line number Diff line number Diff line change
@@ -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

Expand All @@ -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,
Expand Down
11 changes: 9 additions & 2 deletions app/api/endpoints/auth.py
Original file line number Diff line number Diff line change
Expand Up @@ -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)
Expand All @@ -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,
)
Expand All @@ -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
Expand Down
9 changes: 8 additions & 1 deletion app/database/tables.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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)
9 changes: 9 additions & 0 deletions app/ioc.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -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):
Expand Down Expand Up @@ -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):
Expand All @@ -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]
28 changes: 28 additions & 0 deletions app/repositories/revoked_tokens_repository.py
Original file line number Diff line number Diff line change
@@ -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)
14 changes: 14 additions & 0 deletions app/use_cases/check_token_revoked.py
Original file line number Diff line number Diff line change
@@ -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)
19 changes: 19 additions & 0 deletions app/use_cases/revoke_token.py
Original file line number Diff line number Diff line change
@@ -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()
14 changes: 14 additions & 0 deletions docs/adr/0011-logout-revokes-by-jti.md
Original file line number Diff line number Diff line change
@@ -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.
36 changes: 36 additions & 0 deletions migrations/versions/2026-10-10_revoked_tokens.py
Original file line number Diff line number Diff line change
@@ -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 ###
39 changes: 34 additions & 5 deletions tests/api/test_auth_api.py
Original file line number Diff line number Diff line change
Expand Up @@ -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


Expand Down Expand Up @@ -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))
Expand All @@ -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
Expand All @@ -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
Expand All @@ -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]
25 changes: 25 additions & 0 deletions tests/use_cases/test_revoke_token.py
Original file line number Diff line number Diff line change
@@ -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")
Loading