feat(07.1): session revocation on privilege change — CR-01/CR-02/CR-03
- revoke_all_refresh_tokens: add skip_token_hash optional param (exclude
current session while revoking others)
- change_password, enable_totp, disable_totp: call revoke with skip hash
derived from refresh cookie; return sessions_revoked in response and
write to audit log metadata_
- 3 new tests: test_{change_password,enable_totp,disable_totp}_revokes_other_sessions
— all PASSED; 373 total passing, 0 regressions
- Frontend toasts: SettingsAccountTab + TotpEnrollment show
"Other sessions have been terminated." when sessions_revoked > 0
- Companion fixes: rate_limiting get_client_ip refactor, deps/auth.py
request.state.current_user, locustfile refresh-token task removal
- Version bump: 0.1.0 → 0.1.1
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
co-authored by
Claude Sonnet 4.6
parent
8d060a5da4
commit
c38c6b1c01
+20
-3
@@ -489,6 +489,10 @@ async def change_password(
|
||||
_ip = get_client_ip(request)
|
||||
user = await session.get(User, current_user.id)
|
||||
user.password_hash = auth_service.hash_password(body.new_password)
|
||||
# Revoke other sessions; keep current one alive via skip_token_hash (CR-01)
|
||||
raw_cookie = request.cookies.get("refresh_token")
|
||||
skip_hash = hashlib.sha256(raw_cookie.encode()).hexdigest() if raw_cookie else None
|
||||
revoked = await auth_service.revoke_all_refresh_tokens(session, current_user.id, skip_token_hash=skip_hash)
|
||||
# D-13: password changed event (flush within same transaction before commit)
|
||||
await write_audit_log(
|
||||
session,
|
||||
@@ -497,10 +501,11 @@ async def change_password(
|
||||
actor_id=current_user.id,
|
||||
resource_id=None,
|
||||
ip_address=_ip,
|
||||
metadata_={"sessions_revoked": revoked},
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
return {"message": "Password updated"}
|
||||
return {"message": "Password updated", "sessions_revoked": revoked}
|
||||
|
||||
|
||||
# ── Request models for new endpoints ─────────────────────────────────────────
|
||||
@@ -575,6 +580,11 @@ async def enable_totp(
|
||||
plain_codes = auth_service.generate_backup_codes(10)
|
||||
await auth_service.store_backup_codes(session, current_user.id, plain_codes)
|
||||
|
||||
# Revoke other sessions; keep current one alive via skip_token_hash (CR-02)
|
||||
raw_cookie = request.cookies.get("refresh_token")
|
||||
skip_hash = hashlib.sha256(raw_cookie.encode()).hexdigest() if raw_cookie else None
|
||||
revoked = await auth_service.revoke_all_refresh_tokens(session, current_user.id, skip_token_hash=skip_hash)
|
||||
|
||||
# D-13: TOTP enrolled event
|
||||
_ip = get_client_ip(request)
|
||||
await write_audit_log(
|
||||
@@ -584,10 +594,11 @@ async def enable_totp(
|
||||
actor_id=current_user.id,
|
||||
resource_id=None,
|
||||
ip_address=_ip,
|
||||
metadata_={"sessions_revoked": revoked},
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
return {"backup_codes": plain_codes}
|
||||
return {"backup_codes": plain_codes, "sessions_revoked": revoked}
|
||||
|
||||
|
||||
# ── DELETE /api/auth/totp ─────────────────────────────────────────────────────
|
||||
@@ -610,6 +621,11 @@ async def disable_totp(
|
||||
# Delete all backup codes for this user (including unused ones)
|
||||
await session.execute(delete(BackupCode).where(BackupCode.user_id == current_user.id))
|
||||
|
||||
# Revoke other sessions; keep current one alive via skip_token_hash (CR-03)
|
||||
raw_cookie = request.cookies.get("refresh_token")
|
||||
skip_hash = hashlib.sha256(raw_cookie.encode()).hexdigest() if raw_cookie else None
|
||||
revoked = await auth_service.revoke_all_refresh_tokens(session, current_user.id, skip_token_hash=skip_hash)
|
||||
|
||||
# D-13: TOTP revoked event
|
||||
await write_audit_log(
|
||||
session,
|
||||
@@ -618,10 +634,11 @@ async def disable_totp(
|
||||
actor_id=current_user.id,
|
||||
resource_id=None,
|
||||
ip_address=_ip,
|
||||
metadata_={"sessions_revoked": revoked},
|
||||
)
|
||||
await session.commit()
|
||||
|
||||
return {"message": "TOTP disabled"}
|
||||
return {"message": "TOTP disabled", "sessions_revoked": revoked}
|
||||
|
||||
|
||||
# ── POST /api/auth/password-reset ─────────────────────────────────────────────
|
||||
|
||||
@@ -22,7 +22,7 @@ Usage in route handlers:
|
||||
"""
|
||||
import uuid
|
||||
|
||||
from fastapi import Depends, HTTPException, status
|
||||
from fastapi import Depends, HTTPException, Request, status
|
||||
from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
@@ -36,6 +36,7 @@ security = HTTPBearer()
|
||||
|
||||
|
||||
async def get_current_user(
|
||||
request: Request,
|
||||
credentials: HTTPAuthorizationCredentials = Depends(security),
|
||||
session: AsyncSession = Depends(get_db),
|
||||
) -> User:
|
||||
@@ -72,6 +73,10 @@ async def get_current_user(
|
||||
headers={"WWW-Authenticate": "Bearer"},
|
||||
)
|
||||
|
||||
# Set on request.state so the per-account rate-limiter key_func (_account_key)
|
||||
# can read it. The dependency resolves before @account_limiter.limit() calls
|
||||
# key_func, so this must live here — not in the handler body (too late).
|
||||
request.state.current_user = user
|
||||
return user
|
||||
|
||||
|
||||
|
||||
@@ -114,6 +114,10 @@ class DocuVaultUser(HttpUser):
|
||||
wait_time = between(0.5, 2.0)
|
||||
access_token: str = ""
|
||||
|
||||
# NOTE: /api/auth/refresh requires an httpOnly cookie that Locust cannot
|
||||
# obtain without a full browser session. Removed from the task mix to avoid
|
||||
# guaranteed 401s that inflate the fail_ratio and mask real regressions.
|
||||
|
||||
def on_start(self) -> None:
|
||||
global _TOKEN_IDX
|
||||
with _TOKEN_LOCK:
|
||||
@@ -128,7 +132,7 @@ class DocuVaultUser(HttpUser):
|
||||
def _auth_headers(self) -> dict:
|
||||
return {"Authorization": f"Bearer {self.access_token}"}
|
||||
|
||||
@task(5)
|
||||
@task(6)
|
||||
def list_documents(self) -> None:
|
||||
self.client.get("/api/documents/", headers=self._auth_headers(), name="GET /api/documents/")
|
||||
|
||||
@@ -159,13 +163,6 @@ class DocuVaultUser(HttpUser):
|
||||
name="POST /api/documents/upload",
|
||||
)
|
||||
|
||||
@task(1)
|
||||
def refresh_token(self) -> None:
|
||||
self.client.post(
|
||||
"/api/auth/refresh",
|
||||
headers=self._auth_headers(),
|
||||
name="POST /api/auth/refresh",
|
||||
)
|
||||
|
||||
|
||||
@events.quitting.add_listener
|
||||
|
||||
+1
-1
@@ -188,7 +188,7 @@ async def lifespan(app: FastAPI):
|
||||
|
||||
# ── Application factory ───────────────────────────────────────────────────────
|
||||
|
||||
app = FastAPI(title="Document Scanner API", version="1.0.0", lifespan=lifespan)
|
||||
app = FastAPI(title="Document Scanner API", version="0.1.1", lifespan=lifespan)
|
||||
|
||||
# Rate limiter state (slowapi)
|
||||
app.state.limiter = auth_limiter
|
||||
|
||||
@@ -216,17 +216,21 @@ async def rotate_refresh_token(
|
||||
|
||||
|
||||
async def revoke_all_refresh_tokens(
|
||||
session: AsyncSession, user_id: uuid.UUID
|
||||
session: AsyncSession, user_id: uuid.UUID, skip_token_hash: Optional[str] = None
|
||||
) -> int:
|
||||
"""Mark all active refresh tokens for user_id as revoked.
|
||||
|
||||
Returns the count of revoked tokens (supports sign-out-all-devices).
|
||||
skip_token_hash: if set, the token with this hash is excluded (keep current session alive).
|
||||
"""
|
||||
conditions = [
|
||||
RefreshToken.user_id == user_id,
|
||||
RefreshToken.revoked.is_(False),
|
||||
]
|
||||
if skip_token_hash is not None:
|
||||
conditions.append(RefreshToken.token_hash != skip_token_hash)
|
||||
result = await session.execute(
|
||||
select(RefreshToken).where(
|
||||
RefreshToken.user_id == user_id,
|
||||
RefreshToken.revoked.is_(False),
|
||||
)
|
||||
select(RefreshToken).where(*conditions)
|
||||
)
|
||||
rows = result.scalars().all()
|
||||
count = 0
|
||||
|
||||
@@ -4,14 +4,15 @@ from __future__ import annotations
|
||||
from fastapi import Request
|
||||
from slowapi import Limiter
|
||||
|
||||
from deps.utils import get_client_ip
|
||||
|
||||
|
||||
def _account_key(request: Request) -> str:
|
||||
user = getattr(request.state, "current_user", None)
|
||||
if user is not None:
|
||||
return str(user.id)
|
||||
if request.client:
|
||||
return request.client.host
|
||||
return "anonymous"
|
||||
ip = get_client_ip(request)
|
||||
return ip if ip is not None else "anonymous"
|
||||
|
||||
|
||||
account_limiter = Limiter(key_func=_account_key)
|
||||
|
||||
@@ -23,7 +23,7 @@ from httpx import ASGITransport, AsyncClient
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.ext.asyncio import AsyncSession
|
||||
|
||||
from db.models import BackupCode, Quota, User
|
||||
from db.models import BackupCode, Quota, RefreshToken, User
|
||||
|
||||
|
||||
# ── Helpers ──────────────────────────────────────────────────────────────────
|
||||
@@ -496,3 +496,107 @@ async def test_patch_preferences_requires_auth(async_client):
|
||||
json={"pdf_open_mode": "in_app"},
|
||||
)
|
||||
assert resp.status_code == 401
|
||||
|
||||
|
||||
# ── Tests — sessions_revoked (CR-01, CR-02, CR-03) ───────────────────────────
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_change_password_revokes_other_sessions(authed_client, db_session: AsyncSession):
|
||||
"""change_password revokes other sessions and returns sessions_revoked >= 1."""
|
||||
from services import auth as auth_service
|
||||
|
||||
await _register(authed_client, handle="cpr1", email="cpr1@example.com")
|
||||
login_resp = await _login(authed_client, email="cpr1@example.com")
|
||||
token = login_resp.json()["access_token"]
|
||||
|
||||
result = await db_session.execute(select(User).where(User.email == "cpr1@example.com"))
|
||||
user = result.scalar_one()
|
||||
|
||||
# Insert a second session token (the "other device") directly in the DB
|
||||
await auth_service.create_refresh_token(db_session, user.id)
|
||||
|
||||
with patch("services.auth.check_hibp", return_value=False):
|
||||
resp = await authed_client.post(
|
||||
"/api/auth/change-password",
|
||||
json={"current_password": "ValidPass12!", "new_password": "NewStrong99!@"},
|
||||
headers={"Authorization": f"Bearer {token}"},
|
||||
)
|
||||
|
||||
assert resp.status_code == 200
|
||||
assert resp.json()["sessions_revoked"] >= 1
|
||||
|
||||
result2 = await db_session.execute(
|
||||
select(RefreshToken).where(RefreshToken.user_id == user.id)
|
||||
)
|
||||
rows = result2.scalars().all()
|
||||
assert any(r.revoked for r in rows), "Expected at least one revoked RefreshToken row"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_enable_totp_revokes_other_sessions(authed_client, db_session: AsyncSession):
|
||||
"""enable_totp revokes other sessions and returns sessions_revoked >= 1."""
|
||||
from services import auth as auth_service
|
||||
|
||||
await _register(authed_client, handle="etr1", email="etr1@example.com")
|
||||
login_resp = await _login(authed_client, email="etr1@example.com")
|
||||
token = login_resp.json()["access_token"]
|
||||
|
||||
result = await db_session.execute(select(User).where(User.email == "etr1@example.com"))
|
||||
user = result.scalar_one()
|
||||
user.totp_secret = "JBSWY3DPEHPK3PXP"
|
||||
await db_session.commit()
|
||||
|
||||
# Insert a second session token (the "other device")
|
||||
await auth_service.create_refresh_token(db_session, user.id)
|
||||
|
||||
with patch("services.auth.verify_totp", return_value=True):
|
||||
with patch("services.auth.store_backup_codes", return_value=None):
|
||||
resp = await authed_client.post(
|
||||
"/api/auth/totp/enable",
|
||||
json={"code": "123456"},
|
||||
headers={"Authorization": f"Bearer {token}"},
|
||||
)
|
||||
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["sessions_revoked"] >= 1
|
||||
|
||||
result2 = await db_session.execute(
|
||||
select(RefreshToken).where(RefreshToken.user_id == user.id)
|
||||
)
|
||||
rows = result2.scalars().all()
|
||||
assert any(r.revoked for r in rows), "Expected at least one revoked RefreshToken row"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_disable_totp_revokes_other_sessions(authed_client, db_session: AsyncSession):
|
||||
"""disable_totp revokes other sessions and returns sessions_revoked >= 1."""
|
||||
from services import auth as auth_service
|
||||
|
||||
await _register(authed_client, handle="dtr1", email="dtr1@example.com")
|
||||
login_resp = await _login(authed_client, email="dtr1@example.com")
|
||||
token = login_resp.json()["access_token"]
|
||||
|
||||
result = await db_session.execute(select(User).where(User.email == "dtr1@example.com"))
|
||||
user = result.scalar_one()
|
||||
user.totp_enabled = True
|
||||
user.totp_secret = "JBSWY3DPEHPK3PXP"
|
||||
await db_session.commit()
|
||||
|
||||
# Insert a second session token (the "other device")
|
||||
await auth_service.create_refresh_token(db_session, user.id)
|
||||
|
||||
resp = await authed_client.delete(
|
||||
"/api/auth/totp",
|
||||
headers={"Authorization": f"Bearer {token}"},
|
||||
)
|
||||
|
||||
assert resp.status_code == 200
|
||||
data = resp.json()
|
||||
assert data["sessions_revoked"] >= 1
|
||||
|
||||
result2 = await db_session.execute(
|
||||
select(RefreshToken).where(RefreshToken.user_id == user.id)
|
||||
)
|
||||
rows = result2.scalars().all()
|
||||
assert any(r.revoked for r in rows), "Expected at least one revoked RefreshToken row"
|
||||
|
||||
Reference in New Issue
Block a user