Files
home-automation/app/services/totp.py
T

222 lines
7.2 KiB
Python
Raw Normal View History

"""TOTP service: secret management, recovery-code lifecycle, setup/enable/disable (M4-T05).
State-machine summary
---------------------
DISABLED (default)
totp_secret=None, totp_enabled=False, no recovery codes
PENDING (after setup, before enable)
totp_secret=<base32>, totp_enabled=False, recovery codes stored as hashes
- Re-calling setup replaces the pending secret and regenerates recovery codes.
ENABLED (after enable)
totp_secret=<base32>, totp_enabled=True, recovery codes stored as hashes
After disable:
totp_secret=None, totp_enabled=False, recovery codes deleted → back to DISABLED
Recovery-code timing
--------------------
1. setup → generate 10 plaintext codes, persist their Argon2 hashes immediately,
return plaintext to caller (ONLY time).
2. enable → verify 6-digit TOTP code; if ok, flip totp_enabled=True.
Recovery codes are already in the DB from step 1.
3. disable → clear secret, clear enabled flag, delete all recovery codes.
Security note: ``totp_secret`` is stored as plaintext (same posture as other
secrets in this project: protected by file-system permissions on the SQLite
database). Recovery codes are stored as Argon2 hashes only.
"""
from __future__ import annotations
import logging
import secrets
import string
import pyotp
from sqlalchemy import delete, select
from sqlalchemy.orm import Session
from app.models.auth import AuthUser, RecoveryCode
from app.services.auth import hash_password, verify_password
logger = logging.getLogger(__name__)
# Number of recovery codes to generate per setup
_RECOVERY_CODE_COUNT = 10
# Alphabet for each 4-character segment: lowercase letters + digits, no ambiguous chars
_CODE_ALPHABET = string.ascii_lowercase + string.digits
# ---------------------------------------------------------------------------
# Internal helpers
# ---------------------------------------------------------------------------
def _generate_recovery_code() -> str:
"""Return a single recovery code in ``xxxx-xxxx`` format."""
segment = lambda: "".join(secrets.choice(_CODE_ALPHABET) for _ in range(4)) # noqa: E731
return f"{segment()}-{segment()}"
def _verify_totp_code(secret: str, code: str) -> bool:
"""Verify a 6-digit TOTP code against the given base-32 secret (±1 window)."""
return pyotp.TOTP(secret).verify(code, valid_window=1)
def _delete_recovery_codes(db: Session, *, user_id: int) -> None:
"""Delete ALL recovery codes for a user (called on setup re-run and disable)."""
db.execute(delete(RecoveryCode).where(RecoveryCode.user_id == user_id))
# ---------------------------------------------------------------------------
# Public API
# ---------------------------------------------------------------------------
def setup(
db: Session,
*,
user: AuthUser,
issuer: str,
) -> tuple[str, str, list[str]]:
"""Generate a new pending TOTP secret and recovery codes.
The secret is stored in ``user.totp_secret`` (``totp_enabled`` stays
``False``). Any pre-existing pending secret and recovery codes are
replaced atomically.
Returns
-------
(secret, otpauth_uri, plaintext_recovery_codes)
``plaintext_recovery_codes`` are returned **once** here and MUST NOT be
stored or returned elsewhere.
"""
# Generate new TOTP secret
new_secret = pyotp.random_base32()
otpauth_uri = pyotp.TOTP(new_secret).provisioning_uri(
name=user.username,
issuer_name=issuer,
)
# Generate plaintext recovery codes
plaintext_codes = [_generate_recovery_code() for _ in range(_RECOVERY_CODE_COUNT)]
# --- Persist atomically ---
# 1. Delete any old pending recovery codes (idempotent on re-setup)
assert user.id is not None # mypy guard; always set after DB insertion
_delete_recovery_codes(db, user_id=user.id)
# 2. Update the secret (pending — totp_enabled stays False)
user.totp_secret = new_secret
# 3. Persist new recovery codes as Argon2 hashes
for plaintext in plaintext_codes:
db.add(RecoveryCode(user_id=user.id, code_hash=hash_password(plaintext)))
db.commit()
db.refresh(user)
logger.info("TOTP setup (pending) for user '%s'; %d recovery codes generated.", user.username, _RECOVERY_CODE_COUNT)
return new_secret, otpauth_uri, plaintext_codes
def enable(db: Session, *, user: AuthUser, code: str) -> bool:
"""Enable TOTP after the user proves they can generate the correct code.
Parameters
----------
code:
The current 6-digit TOTP code from the user's authenticator app.
Returns
-------
True on success; False if the code is invalid or no secret is pending.
"""
if not user.totp_secret:
logger.info("TOTP enable rejected for '%s': no pending secret.", user.username)
return False
if not _verify_totp_code(user.totp_secret, code):
logger.info("TOTP enable rejected for '%s': wrong code.", user.username)
return False
user.totp_enabled = True
db.commit()
db.refresh(user)
logger.info("TOTP enabled for user '%s'.", user.username)
return True
def disable(
db: Session,
*,
user: AuthUser,
password: str | None = None,
code: str | None = None,
) -> bool:
"""Disable TOTP. Caller must supply exactly one of ``password`` or ``code``.
On success: ``totp_enabled=False``, ``totp_secret=None``, all recovery
codes deleted.
Returns
-------
True on success; False if neither credential is valid.
"""
if not password and not code:
logger.info("TOTP disable rejected for '%s': no credential provided.", user.username)
return False
if password:
if not verify_password(password, user.password_hash):
logger.info("TOTP disable rejected for '%s': wrong password.", user.username)
return False
elif code:
if not user.totp_secret or not _verify_totp_code(user.totp_secret, code):
logger.info("TOTP disable rejected for '%s': wrong TOTP code.", user.username)
return False
# Clear TOTP state
assert user.id is not None
_delete_recovery_codes(db, user_id=user.id)
user.totp_enabled = False
user.totp_secret = None
db.commit()
db.refresh(user)
logger.info("TOTP disabled for user '%s'.", user.username)
return True
def verify_recovery_code(db: Session, *, user: AuthUser, code: str) -> bool:
"""Verify and consume a one-time recovery code.
Finds the first unused recovery code whose hash matches ``code``, marks it
as consumed (sets ``used_at``), and returns ``True``. Returns ``False`` if
no matching unused code exists.
This function is provided for T06 (login two-factor) but lives here so the
TOTP service owns all recovery-code logic.
"""
from datetime import UTC, datetime
unused = db.execute(
select(RecoveryCode).where(
RecoveryCode.user_id == user.id,
RecoveryCode.used_at.is_(None),
)
).scalars().all()
for rc in unused:
if verify_password(code, rc.code_hash):
rc.used_at = datetime.now(UTC)
db.commit()
logger.info(
"Recovery code consumed for user '%s' (id=%d).", user.username, rc.id
)
return True
return False