"""Issuing, spending and withdrawing refresh tokens.""" import hashlib import secrets from datetime import datetime, timedelta from sqlalchemy.orm import Session from app.models.refresh_token import RefreshToken from app.models.user import User #: Long enough that an app is not signing people out every month; short enough #: that a token forgotten on a lost phone stops working within a season. LIFETIME = timedelta(days=60) #: 256 bits from the system generator. Not a JWT: there is nothing to read in #: it, and everything about it is settled by the row it points at. BYTES = 32 #: One person, one device, one session. A cap stops a buggy client filling the #: table with families nobody will ever spend. MAX_FAMILIES = 20 def _hash(token: str) -> str: return hashlib.sha256(token.encode()).hexdigest() def issue(db: Session, user: User, *, label: str | None = None, family: str | None = None, ip: str | None = None) -> str: """Mint a token, return it once. Only its hash is kept.""" # Trimmed before the new row exists, not after: querying with a pending # insert in the session makes SQLAlchemy reload it mid-flush. if family is None: _trim(db, user) token = secrets.token_urlsafe(BYTES) db.add(RefreshToken( user_id=user.id, token_hash=_hash(token), family=family or secrets.token_hex(8), label=(label or "")[:120] or None, expires_at=datetime.utcnow() + LIFETIME, last_ip=(ip or "")[:64] or None, )) db.commit() return token def _trim(db: Session, user: User) -> None: """Keep the newest families; end the rest.""" families = [row.family for row in db.query(RefreshToken).filter( RefreshToken.user_id == user.id, RefreshToken.revoked_at.is_(None), ).order_by(RefreshToken.created_at.desc()).all()] seen: list[str] = [] for name in families: if name not in seen: seen.append(name) for name in seen[MAX_FAMILIES - 1:]: revoke_family(db, user, name, commit=False) def spend(db: Session, token: str, *, ip: str | None = None) -> tuple[User, str] | None: """Exchange a token for its user and a fresh token, or None if it is no good. A token that has already been spent ends its whole family. Either somebody copied it, or a client is replaying — and there is no way to tell which from here, so the safe reading is the unsafe one. """ row = db.query(RefreshToken).filter(RefreshToken.token_hash == _hash(token)).first() if row is None: return None user = db.get(User, row.user_id) if user is None: return None if row.used_at is not None: revoke_family(db, user, row.family) return None if row.revoked_at is not None or row.expires_at <= datetime.utcnow(): return None row.used_at = datetime.utcnow() if ip: row.last_ip = ip[:64] db.commit() return user, issue(db, user, label=row.label, family=row.family, ip=ip) def revoke_family(db: Session, user: User, family: str, *, commit: bool = True) -> int: count = db.query(RefreshToken).filter( RefreshToken.user_id == user.id, RefreshToken.family == family, RefreshToken.revoked_at.is_(None), ).update({"revoked_at": datetime.utcnow()}, synchronize_session=False) if commit: db.commit() return count def revoke_one(db: Session, token: str) -> bool: """End the session a token belongs to. Used by sign-out.""" row = db.query(RefreshToken).filter(RefreshToken.token_hash == _hash(token)).first() if row is None: return False user = db.get(User, row.user_id) if user is None: return False revoke_family(db, user, row.family) return True def revoke_all(db: Session, user: User) -> int: """Sign out everywhere. What a person wants after losing a phone.""" count = db.query(RefreshToken).filter( RefreshToken.user_id == user.id, RefreshToken.revoked_at.is_(None), ).update({"revoked_at": datetime.utcnow()}, synchronize_session=False) db.commit() return count def sessions(db: Session, user: User) -> list[dict]: """Where this person is signed in, one row per family.""" rows = db.query(RefreshToken).filter( RefreshToken.user_id == user.id, RefreshToken.revoked_at.is_(None), RefreshToken.expires_at > datetime.utcnow(), ).order_by(RefreshToken.created_at.desc()).all() seen: dict[str, dict] = {} for row in rows: # The newest token in a family describes the session now. seen.setdefault(row.family, { "family": row.family, "label": row.label, "started_at": row.created_at, "expires_at": row.expires_at, "last_ip": row.last_ip, "current": False, }) return list(seen.values()) def purge(db: Session) -> int: """Drop what can never be spent again. Called by the nightly sweep.""" cutoff = datetime.utcnow() - timedelta(days=7) return db.query(RefreshToken).filter( (RefreshToken.expires_at < cutoff) | (RefreshToken.revoked_at < cutoff), ).delete(synchronize_session=False)