diff --git a/.env.example b/.env.example index b42109c..9a6e273 100644 --- a/.env.example +++ b/.env.example @@ -14,7 +14,12 @@ DB_KEEPALIVES_COUNT=5 # Generate a value with: openssl rand -hex 32 SECRET_KEY=replace-with-a-random-secret -ACCESS_TOKEN_EXPIRE_MINUTES=60 +ACCESS_TOKEN_EXPIRE_MINUTES=15 +REFRESH_TOKEN_EXPIRE_DAYS=30 +REFRESH_TOKEN_IDLE_DAYS=7 +REFRESH_COOKIE_NAME=htv_refresh +REFRESH_COOKIE_SECURE=false +REFRESH_COOKIE_SAMESITE=lax ACTIVATION_TOKEN_EXPIRE_MINUTES=60 PASSWORD_RESET_TOKEN_EXPIRE_MINUTES=15 PASSWORD_RESET_COOLDOWN_MINUTES=15 diff --git a/.env.prod.example b/.env.prod.example index 4826da8..035f2e0 100644 --- a/.env.prod.example +++ b/.env.prod.example @@ -8,7 +8,12 @@ POSTGRES_DB=postgres DATABASE_URL=postgresql://postgres:replace-with-the-same-password@db:5432/postgres SECRET_KEY=replace-with-output-of-openssl-rand-hex-32 -ACCESS_TOKEN_EXPIRE_MINUTES=60 +ACCESS_TOKEN_EXPIRE_MINUTES=15 +REFRESH_TOKEN_EXPIRE_DAYS=30 +REFRESH_TOKEN_IDLE_DAYS=7 +REFRESH_COOKIE_NAME=htv_refresh +REFRESH_COOKIE_SECURE=true +REFRESH_COOKIE_SAMESITE=lax ACTIVATION_TOKEN_EXPIRE_MINUTES=60 PASSWORD_RESET_TOKEN_EXPIRE_MINUTES=15 PASSWORD_RESET_COOLDOWN_MINUTES=15 diff --git a/alembic/env.py b/alembic/env.py index 2f15d4e..94719c4 100644 --- a/alembic/env.py +++ b/alembic/env.py @@ -6,7 +6,15 @@ from sqlmodel import SQLModel from alembic import context -from app.models import bulk_email, food_tracking, forms, judging, meal, user # noqa: F401 +from app.models import ( # noqa: F401 + bulk_email, + food_tracking, + forms, + judging, + meal, + refresh_session, + user, +) # this is the Alembic Config object, which provides # access to the values within the .ini file in use. diff --git a/alembic/versions/e84b6a92c1d4_add_refresh_sessions.py b/alembic/versions/e84b6a92c1d4_add_refresh_sessions.py new file mode 100644 index 0000000..671b548 --- /dev/null +++ b/alembic/versions/e84b6a92c1d4_add_refresh_sessions.py @@ -0,0 +1,71 @@ +"""add refresh sessions + +Revision ID: e84b6a92c1d4 +Revises: a31f0e8c4d12 +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op +from sqlalchemy.dialects import postgresql + +revision: str = "e84b6a92c1d4" +down_revision: str | None = "a31f0e8c4d12" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +def upgrade() -> None: + op.create_table( + "account_refresh_session", + sa.Column("session_id", postgresql.UUID(as_uuid=True), nullable=False), + sa.Column("user_id", postgresql.UUID(as_uuid=True), nullable=False), + sa.Column("family_id", postgresql.UUID(as_uuid=True), nullable=False), + sa.Column("token_hash", sa.String(length=64), nullable=False), + sa.Column("token_version", sa.Integer(), nullable=False), + sa.Column("created_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("last_used_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("expires_at", sa.DateTime(timezone=True), nullable=False), + sa.Column("revoked_at", sa.DateTime(timezone=True), nullable=True), + sa.Column("replaced_by_id", postgresql.UUID(as_uuid=True), nullable=True), + sa.ForeignKeyConstraint( + ["replaced_by_id"], + ["account_refresh_session.session_id"], + ondelete="SET NULL", + ), + sa.ForeignKeyConstraint(["user_id"], ["account_user.uid"], ondelete="CASCADE"), + sa.PrimaryKeyConstraint("session_id"), + ) + op.create_index( + op.f("ix_account_refresh_session_family_id"), + "account_refresh_session", + ["family_id"], + ) + op.create_index( + op.f("ix_account_refresh_session_token_hash"), + "account_refresh_session", + ["token_hash"], + unique=True, + ) + op.create_index( + op.f("ix_account_refresh_session_user_id"), + "account_refresh_session", + ["user_id"], + ) + + +def downgrade() -> None: + op.drop_index( + op.f("ix_account_refresh_session_user_id"), + table_name="account_refresh_session", + ) + op.drop_index( + op.f("ix_account_refresh_session_token_hash"), + table_name="account_refresh_session", + ) + op.drop_index( + op.f("ix_account_refresh_session_family_id"), + table_name="account_refresh_session", + ) + op.drop_table("account_refresh_session") diff --git a/app/config.py b/app/config.py index 4179f5c..551cc91 100644 --- a/app/config.py +++ b/app/config.py @@ -2,7 +2,7 @@ import os from datetime import datetime from functools import lru_cache -from typing import Annotated, Any +from typing import Annotated, Any, Literal from pydantic import Field, field_validator from pydantic_settings import BaseSettings, NoDecode, SettingsConfigDict @@ -30,7 +30,12 @@ class Settings(BaseSettings): SECRET_KEY: str = "" ALGORITHM: str = "HS256" - ACCESS_TOKEN_EXPIRE_MINUTES: int = 60 + ACCESS_TOKEN_EXPIRE_MINUTES: int = 15 + REFRESH_TOKEN_EXPIRE_DAYS: int = 30 + REFRESH_TOKEN_IDLE_DAYS: int = 7 + REFRESH_COOKIE_NAME: str = "htv_refresh" + REFRESH_COOKIE_SECURE: bool = True + REFRESH_COOKIE_SAMESITE: Literal["lax", "strict", "none"] = "lax" ACTIVATION_TOKEN_EXPIRE_MINUTES: int = 60 PASSWORD_RESET_TOKEN_EXPIRE_MINUTES: int = 15 PASSWORD_RESET_COOLDOWN_MINUTES: int = 15 diff --git a/app/models/refresh_session.py b/app/models/refresh_session.py new file mode 100644 index 0000000..90a0fca --- /dev/null +++ b/app/models/refresh_session.py @@ -0,0 +1,42 @@ +import uuid +from datetime import datetime, timezone +from typing import ClassVar + +from sqlmodel import Column, DateTime, Field, SQLModel, String + + +def utc_now() -> datetime: + return datetime.now(timezone.utc) + + +class RefreshSession(SQLModel, table=True): + __tablename__: ClassVar[str] = "account_refresh_session" + + session_id: uuid.UUID = Field(default_factory=uuid.uuid4, primary_key=True) + user_id: uuid.UUID = Field( + foreign_key="account_user.uid", ondelete="CASCADE", index=True + ) + family_id: uuid.UUID = Field(default_factory=uuid.uuid4, index=True) + token_hash: str = Field( + sa_column=Column(String(64), nullable=False, unique=True, index=True) + ) + token_version: int = Field(nullable=False) + created_at: datetime = Field( + default_factory=utc_now, + sa_column=Column(DateTime(timezone=True), nullable=False), + ) + last_used_at: datetime = Field( + default_factory=utc_now, + sa_column=Column(DateTime(timezone=True), nullable=False), + ) + expires_at: datetime = Field( + sa_column=Column(DateTime(timezone=True), nullable=False) + ) + revoked_at: datetime | None = Field( + default=None, sa_column=Column(DateTime(timezone=True), nullable=True) + ) + replaced_by_id: uuid.UUID | None = Field( + default=None, + foreign_key="account_refresh_session.session_id", + ondelete="SET NULL", + ) diff --git a/app/routers/account.py b/app/routers/account.py index f28e106..a785032 100644 --- a/app/routers/account.py +++ b/app/routers/account.py @@ -3,14 +3,23 @@ from typing import Annotated import bcrypt -from fastapi import APIRouter, BackgroundTasks, Depends, HTTPException, Query, status -from fastapi.responses import Response +from fastapi import ( + APIRouter, + BackgroundTasks, + Depends, + HTTPException, + Query, + Request, + Response, + status, +) from fastapi.security import OAuth2PasswordRequestForm from sqlmodel import select from sqlalchemy.exc import IntegrityError, SQLAlchemyError from app.config import AppConfig, SecurityConfig from app.core.db import SessionDep +from app.core.errors import ServiceError from app.core.orm import eager_load from app.models.constants import ( EmailMessage, @@ -29,6 +38,11 @@ decode_token, scopes_for_user, ) +from app.services.refresh_sessions import ( + create_refresh_session, + revoke_refresh_session, + rotate_refresh_session, +) from app.services.email import ( send_activation_email_in_background, send_email, @@ -54,11 +68,34 @@ def _invalid_login() -> HTTPException: ) +def _set_refresh_cookie(response: Response, token: str) -> None: + response.set_cookie( + key=SecurityConfig.REFRESH_COOKIE_NAME, + value=token, + max_age=SecurityConfig.REFRESH_TOKEN_EXPIRE_DAYS * 24 * 60 * 60, + path="/api/account/tokens", + secure=SecurityConfig.REFRESH_COOKIE_SECURE, + httponly=True, + samesite=SecurityConfig.REFRESH_COOKIE_SAMESITE, + ) + + +def _clear_refresh_cookie(response: Response) -> None: + response.delete_cookie( + key=SecurityConfig.REFRESH_COOKIE_NAME, + path="/api/account/tokens", + secure=SecurityConfig.REFRESH_COOKIE_SECURE, + httponly=True, + samesite=SecurityConfig.REFRESH_COOKIE_SAMESITE, + ) + + @router.post("/sessions") def login( form_data: Annotated[OAuth2PasswordRequestForm, Depends()], session: SessionDep, background_tasks: BackgroundTasks, + response: Response, ) -> Token: statement = select(AccountUser).where(AccountUser.email == form_data.username) selected_user = session.exec(statement).first() @@ -108,6 +145,9 @@ def login( scopes = scopes_for_user(selected_user) access_token_expires = timedelta(minutes=SecurityConfig.ACCESS_TOKEN_EXPIRE_MINUTES) access_token = create_user_access_token(selected_user, scopes, access_token_expires) + _, refresh_token = create_refresh_session(session, selected_user) + session.commit() + _set_refresh_cookie(response, refresh_token) return Token(access_token=access_token, token_type="bearer") @@ -268,15 +308,41 @@ def activate(user: UserUpdate, session: SessionDep) -> bool: @router.post("/tokens") def refresh( - current_user: Annotated[AccountUser, Depends(get_current_user)], + request: Request, + response: Response, + session: SessionDep, ) -> Token: + refresh_token = request.cookies.get(SecurityConfig.REFRESH_COOKIE_NAME) + if not refresh_token: + raise HTTPException( + status_code=status.HTTP_401_UNAUTHORIZED, + detail="Refresh token required", + ) + try: + current_user, replacement_token = rotate_refresh_session(session, refresh_token) + except ServiceError as error: + _clear_refresh_cookie(response) + raise HTTPException( + status_code=error.status_code, + detail=error.detail, + headers={"Set-Cookie": response.headers["set-cookie"]}, + ) from error access_token_expires = timedelta(minutes=SecurityConfig.ACCESS_TOKEN_EXPIRE_MINUTES) access_token = create_user_access_token( current_user, scopes_for_user(current_user), access_token_expires ) + _set_refresh_cookie(response, replacement_token) return Token(access_token=access_token, token_type="bearer") +@router.delete("/tokens", status_code=status.HTTP_204_NO_CONTENT) +def logout(request: Request, response: Response, session: SessionDep) -> None: + revoke_refresh_session( + session, request.cookies.get(SecurityConfig.REFRESH_COOKIE_NAME) + ) + _clear_refresh_cookie(response) + + @router.get("/apple-wallet/{application_id}") def apple_wallet( application_id: str, diff --git a/app/services/refresh_sessions.py b/app/services/refresh_sessions.py new file mode 100644 index 0000000..12b26e3 --- /dev/null +++ b/app/services/refresh_sessions.py @@ -0,0 +1,106 @@ +import hashlib +import secrets +import uuid +from datetime import datetime, timedelta, timezone + +from sqlalchemy import update +from sqlmodel import Session, col, select + +from app.config import SecurityConfig +from app.core.errors import ServiceError +from app.models.refresh_session import RefreshSession +from app.models.user import AccountUser + + +def _hash_token(token: str) -> str: + return hashlib.sha256(token.encode("utf-8")).hexdigest() + + +def _new_token() -> str: + return secrets.token_urlsafe(48) + + +def create_refresh_session( + session: Session, user: AccountUser +) -> tuple[RefreshSession, str]: + now = datetime.now(timezone.utc) + raw_token = _new_token() + refresh_session = RefreshSession( + user_id=user.uid, + token_hash=_hash_token(raw_token), + token_version=user.token_version, + created_at=now, + last_used_at=now, + expires_at=now + timedelta(days=SecurityConfig.REFRESH_TOKEN_EXPIRE_DAYS), + ) + session.add(refresh_session) + return refresh_session, raw_token + + +def _revoke_family( + session: Session, family_id: uuid.UUID, revoked_at: datetime +) -> None: + session.exec( + update(RefreshSession) + .where( + RefreshSession.family_id == family_id, + col(RefreshSession.revoked_at).is_(None), + ) + .values(revoked_at=revoked_at) + ) + + +def rotate_refresh_session(session: Session, raw_token: str) -> tuple[AccountUser, str]: + now = datetime.now(timezone.utc) + current = session.exec( + select(RefreshSession) + .where(RefreshSession.token_hash == _hash_token(raw_token)) + .with_for_update() + ).first() + if current is None: + raise ServiceError(status_code=401, detail="Invalid refresh token") + + if current.revoked_at is not None: + _revoke_family(session, current.family_id, now) + session.commit() + raise ServiceError(status_code=401, detail="Refresh token reuse detected") + + user = session.get(AccountUser, current.user_id) + idle_deadline = current.last_used_at + timedelta( + days=SecurityConfig.REFRESH_TOKEN_IDLE_DAYS + ) + if ( + user is None + or not user.is_active + or user.token_version != current.token_version + or now >= current.expires_at + or now >= idle_deadline + ): + _revoke_family(session, current.family_id, now) + session.commit() + raise ServiceError(status_code=401, detail="Refresh session expired") + + replacement, replacement_token = create_refresh_session(session, user) + replacement.family_id = current.family_id + session.flush() + current.last_used_at = now + current.revoked_at = now + current.replaced_by_id = replacement.session_id + session.add(current) + session.add(replacement) + session.commit() + return user, replacement_token + + +def revoke_refresh_session(session: Session, raw_token: str | None) -> None: + if not raw_token: + return + current = session.exec( + select(RefreshSession).where( + RefreshSession.token_hash == _hash_token(raw_token) + ) + ).first() + if current is None: + return + _revoke_family(session, current.family_id, datetime.now(timezone.utc)) + session.commit() diff --git a/docker-compose.e2e.yml b/docker-compose.e2e.yml index e63bc9e..4abb8b1 100644 --- a/docker-compose.e2e.yml +++ b/docker-compose.e2e.yml @@ -45,6 +45,7 @@ services: environment: DATABASE_URL: postgresql://postgres:postgres@e2e-db:5432/hack_the_back_e2e SECRET_KEY: e2e-only-secret-key-not-for-production + REFRESH_COOKIE_SECURE: "false" POSTMARK_KEY: e2e-postmark-key POSTMARK_URL: http://e2e-mail:8080/email FRONTEND_URL: http://frontend.e2e.test diff --git a/tests/e2e/test_00_seeding.py b/tests/e2e/test_00_seeding.py index e39727c..31453e2 100644 --- a/tests/e2e/test_00_seeding.py +++ b/tests/e2e/test_00_seeding.py @@ -15,7 +15,7 @@ def test_seed_contract_matches_source_data(client, admin_headers): - assert db_query("SELECT version_num FROM alembic_version") == ["a31f0e8c4d12"] + assert db_query("SELECT version_num FROM alembic_version") == ["e84b6a92c1d4"] expected_questions = json.loads( (ROOT / "app/data/form_questions.json").read_text(encoding="utf-8") ) diff --git a/tests/e2e/test_account_and_access.py b/tests/e2e/test_account_and_access.py index d077fd5..78b904d 100644 --- a/tests/e2e/test_account_and_access.py +++ b/tests/e2e/test_account_and_access.py @@ -2,6 +2,7 @@ from datetime import timedelta import httpx +import jwt from .conftest import MAIL_URL, PASSWORD, db_query, token @@ -56,12 +57,20 @@ def test_signup_activation_login_refresh_and_me(client, unique_email): data={"username": unique_email, "password": PASSWORD}, ) assert login.status_code == 200 + access_claims = jwt.decode( + login.json()["access_token"], options={"verify_signature": False} + ) + assert access_claims["exp"] - access_claims["iat"] == 15 * 60 + cookie_header = login.headers["set-cookie"] + assert "HttpOnly" in cookie_header + assert "SameSite=lax" in cookie_header + assert "Path=/api/account/tokens" in cookie_header headers = {"Authorization": f"Bearer {login.json()['access_token']}"} me = client.get("/api/account/me", headers=headers) assert me.status_code == 200 assert me.json()["email"] == unique_email - refresh = client.post("/api/account/tokens", headers=headers) + refresh = client.post("/api/account/tokens") assert refresh.status_code == 200 assert refresh.json()["access_token"] @@ -219,13 +228,8 @@ def test_account_validation_and_failure_paths(client, active_hacker): ).status_code == 401 ) - assert ( - client.post( - "/api/account/tokens", - headers={"Authorization": f"Bearer {wrong_scope}"}, - ).status_code - == 401 - ) + client.cookies.clear() + assert client.post("/api/account/tokens").status_code == 401 assert ( client.get( @@ -235,6 +239,32 @@ def test_account_validation_and_failure_paths(client, active_hacker): ) +def test_refresh_rotation_reuse_detection_and_logout(client, active_hacker): + cookie_name = "htv_refresh" + original = client.cookies.get(cookie_name) + assert original + + refreshed = client.post("/api/account/tokens") + assert refreshed.status_code == 200 + replacement = client.cookies.get(cookie_name) + assert replacement and replacement != original + + client.cookies.set(cookie_name, original, path="/api/account/tokens") + replay = client.post("/api/account/tokens") + assert replay.status_code == 401 + + client.cookies.set(cookie_name, replacement, path="/api/account/tokens") + assert client.post("/api/account/tokens").status_code == 401 + + login = client.post( + "/api/account/sessions", + data={"username": active_hacker["email"], "password": PASSWORD}, + ) + assert login.status_code == 200 + assert client.delete("/api/account/tokens").status_code == 204 + assert client.post("/api/account/tokens").status_code == 401 + + def test_inactive_account_resends_are_generic_and_throttled(client, unique_email): payload = { "first_name": "Inactive", diff --git a/tests/unit/test_refresh_sessions.py b/tests/unit/test_refresh_sessions.py new file mode 100644 index 0000000..a1af839 --- /dev/null +++ b/tests/unit/test_refresh_sessions.py @@ -0,0 +1,144 @@ +from datetime import datetime, timedelta, timezone +from types import SimpleNamespace +from unittest.mock import MagicMock, patch +from uuid import uuid4 + +import pytest + +from app.core.errors import ServiceError +from app.models.refresh_session import RefreshSession +from app.services import refresh_sessions + + +def result(*, first=None): + return SimpleNamespace(first=lambda: first) + + +def user(): + return SimpleNamespace(uid=uuid4(), token_version=3, is_active=True) + + +def stored_session(account, **updates): + now = datetime.now(timezone.utc) + values = { + "user_id": account.uid, + "family_id": uuid4(), + "token_hash": "a" * 64, + "token_version": account.token_version, + "created_at": now, + "last_used_at": now, + "expires_at": now + timedelta(days=30), + } + values.update(updates) + return RefreshSession(**values) + + +def test_create_refresh_session_hashes_token_and_sets_expiry(): + session = MagicMock() + account = user() + + with patch.object(refresh_sessions, "_new_token", return_value="raw-secret"): + created, raw_token = refresh_sessions.create_refresh_session(session, account) + + assert raw_token == "raw-secret" + assert created.token_hash != raw_token + assert created.token_hash == refresh_sessions._hash_token(raw_token) + assert created.user_id == account.uid + assert created.token_version == account.token_version + assert created.expires_at > created.created_at + session.add.assert_called_once_with(created) + + +def test_rotate_rejects_unknown_token(): + session = MagicMock() + session.exec.return_value = result(first=None) + + with pytest.raises(ServiceError, match="Invalid refresh token"): + refresh_sessions.rotate_refresh_session(session, "missing") + + +def test_rotate_detects_reuse_and_revokes_family(): + session = MagicMock() + account = user() + current = stored_session( + account, revoked_at=datetime.now(timezone.utc) - timedelta(seconds=1) + ) + session.exec.side_effect = [result(first=current), MagicMock()] + + with pytest.raises(ServiceError, match="reuse detected"): + refresh_sessions.rotate_refresh_session(session, "replayed") + + session.commit.assert_called_once() + + +@pytest.mark.parametrize( + "change", + [ + {"user_missing": True}, + {"is_active": False}, + {"token_version": 4}, + {"expires_at": datetime.now(timezone.utc) - timedelta(seconds=1)}, + {"last_used_at": datetime.now(timezone.utc) - timedelta(days=8)}, + ], +) +def test_rotate_rejects_invalid_or_expired_session(change): + session = MagicMock() + account = user() + current_updates = { + key: value + for key, value in change.items() + if key in {"expires_at", "last_used_at"} + } + current = stored_session(account, **current_updates) + if "is_active" in change: + account.is_active = change["is_active"] + if "token_version" in change: + account.token_version = change["token_version"] + session.exec.side_effect = [result(first=current), MagicMock()] + session.get.return_value = None if change.get("user_missing") else account + + with pytest.raises(ServiceError, match="expired"): + refresh_sessions.rotate_refresh_session(session, "expired") + + session.commit.assert_called_once() + + +def test_rotate_replaces_valid_token(): + session = MagicMock() + account = user() + current = stored_session(account) + replacement = stored_session(account, family_id=uuid4()) + session.exec.return_value = result(first=current) + session.get.return_value = account + + with patch.object( + refresh_sessions, + "create_refresh_session", + return_value=(replacement, "replacement-secret"), + ): + returned_user, raw_token = refresh_sessions.rotate_refresh_session( + session, "current-secret" + ) + + assert returned_user is account + assert raw_token == "replacement-secret" + assert replacement.family_id == current.family_id + assert current.revoked_at is not None + assert current.replaced_by_id == replacement.session_id + session.commit.assert_called_once() + + +def test_revoke_handles_absent_unknown_and_valid_tokens(): + session = MagicMock() + refresh_sessions.revoke_refresh_session(session, None) + session.exec.assert_not_called() + + session.exec.return_value = result(first=None) + refresh_sessions.revoke_refresh_session(session, "unknown") + session.commit.assert_not_called() + + account = user() + current = stored_session(account) + session.exec.side_effect = [result(first=current), MagicMock()] + refresh_sessions.revoke_refresh_session(session, "valid") + session.commit.assert_called_once()