"""Tests for M4-T05: TOTP service + setup/enable/disable/status API.""" from __future__ import annotations import pyotp import pytest from fastapi.testclient import TestClient from sqlalchemy import create_engine, select from sqlalchemy.orm import Session from app.db import reset_db_caches from app.models.auth import AuthUser, RecoveryCode from app.services.auth import verify_password # --------------------------------------------------------------------------- # Helpers # --------------------------------------------------------------------------- def _login(client: TestClient) -> str: """Log in and return the CSRF token from the response.""" resp = client.post( "/api/auth/login", json={"username": "admin", "password": "test-password"}, ) assert resp.status_code == 200, resp.text return resp.json()["csrf_token"] def _setup(client: TestClient, csrf: str) -> dict: resp = client.post( "/api/auth/totp/setup", headers={"X-CSRF-Token": csrf}, ) return resp def _enable(client: TestClient, csrf: str, code: str): # noqa: ANN201 return client.post( "/api/auth/totp/enable", json={"code": code}, headers={"X-CSRF-Token": csrf}, ) def _disable(client: TestClient, csrf: str, *, password: str | None = None, code: str | None = None): body: dict = {} if password is not None: body["password"] = password if code is not None: body["code"] = code return client.post( "/api/auth/totp/disable", json=body, headers={"X-CSRF-Token": csrf}, ) def _status(client: TestClient) -> dict: resp = client.get("/api/auth/totp") assert resp.status_code == 200, resp.text return resp.json() # --------------------------------------------------------------------------- # Fixtures — reuse conftest app/client via the session-aware totp_client fixture # --------------------------------------------------------------------------- @pytest.fixture def totp_client(auth_database): """TestClient + a direct SQLAlchemy engine so tests can inspect DB state.""" from app.main import create_app app_url = auth_database["app_url"] engine = create_engine(app_url, connect_args={"check_same_thread": False}) fastapi_app = create_app() with TestClient(fastapi_app) as tc: yield tc, engine engine.dispose() reset_db_caches() # --------------------------------------------------------------------------- # POST /api/auth/totp/setup # --------------------------------------------------------------------------- def test_setup_returns_secret_uri_and_recovery_codes(totp_client): client, engine = totp_client csrf = _login(client) resp = _setup(client, csrf) assert resp.status_code == 200, resp.text body = resp.json() assert "secret" in body assert "otpauth_uri" in body assert "recovery_codes" in body assert len(body["recovery_codes"]) == 10 # Each code must match xxxx-xxxx pattern for rc in body["recovery_codes"]: parts = rc.split("-") assert len(parts) == 2 assert all(len(p) == 4 for p in parts) # URI must start with otpauth://totp/ assert body["otpauth_uri"].startswith("otpauth://totp/") # Secret must be a valid base32 string (pyotp will accept it) assert len(body["secret"]) >= 16 def test_setup_does_not_enable_totp(totp_client): """After setup, totp_enabled must still be False.""" client, engine = totp_client csrf = _login(client) _setup(client, csrf) status_body = _status(client) assert status_body["enabled"] is False def test_setup_persists_secret_and_code_hashes(totp_client): """After setup, DB has totp_secret set and recovery code hashes stored.""" client, engine = totp_client csrf = _login(client) resp = _setup(client, csrf) body = resp.json() with Session(engine) as db: user = db.scalar(select(AuthUser).where(AuthUser.username == "admin")) assert user is not None # Secret is stored assert user.totp_secret == body["secret"] # Enabled is still False assert user.totp_enabled is False # Recovery codes are stored as hashes codes = db.execute(select(RecoveryCode).where(RecoveryCode.user_id == user.id)).scalars().all() assert len(codes) == 10 # Hashes are not the same as plaintext for i, rc in enumerate(codes): assert rc.code_hash != body["recovery_codes"][i] # But they must verify correctly with Argon2 assert verify_password(body["recovery_codes"][i], rc.code_hash) def test_setup_twice_replaces_pending_state(totp_client): """Re-running setup clears old pending secret + codes and issues new ones.""" client, engine = totp_client csrf = _login(client) first = _setup(client, csrf).json() second = _setup(client, csrf).json() # Different secret each time assert first["secret"] != second["secret"] # Different codes assert first["recovery_codes"] != second["recovery_codes"] # DB has exactly 10 codes (not 20) with Session(engine) as db: user = db.scalar(select(AuthUser).where(AuthUser.username == "admin")) codes = db.execute( select(RecoveryCode).where(RecoveryCode.user_id == user.id) ).scalars().all() assert len(codes) == 10 # --------------------------------------------------------------------------- # POST /api/auth/totp/enable # --------------------------------------------------------------------------- def test_enable_with_wrong_code_fails(totp_client): client, engine = totp_client csrf = _login(client) _setup(client, csrf) resp = _enable(client, csrf, "000000") assert resp.status_code == 400 def test_enable_with_wrong_code_does_not_enable(totp_client): client, engine = totp_client csrf = _login(client) _setup(client, csrf) _enable(client, csrf, "000000") assert _status(client)["enabled"] is False def test_enable_with_correct_code_succeeds(totp_client): """Using pyotp.TOTP(secret).now() to generate the correct code.""" client, engine = totp_client csrf = _login(client) setup_body = _setup(client, csrf).json() secret = setup_body["secret"] correct_code = pyotp.TOTP(secret).now() resp = _enable(client, csrf, correct_code) assert resp.status_code == 204, resp.text def test_enable_with_correct_code_sets_totp_enabled(totp_client): client, engine = totp_client csrf = _login(client) setup_body = _setup(client, csrf).json() secret = setup_body["secret"] _enable(client, csrf, pyotp.TOTP(secret).now()) assert _status(client)["enabled"] is True def test_enable_persists_recovery_code_hashes(totp_client): """Recovery codes are already hashed in DB after enable (not plaintext).""" client, engine = totp_client csrf = _login(client) setup_body = _setup(client, csrf).json() secret = setup_body["secret"] _enable(client, csrf, pyotp.TOTP(secret).now()) with Session(engine) as db: user = db.scalar(select(AuthUser).where(AuthUser.username == "admin")) codes = db.execute(select(RecoveryCode).where(RecoveryCode.user_id == user.id)).scalars().all() assert len(codes) == 10 for i, rc in enumerate(codes): plaintext = setup_body["recovery_codes"][i] # code_hash must NOT be the plaintext assert rc.code_hash != plaintext # But verify_password must return True (Argon2 hash matches) assert verify_password(plaintext, rc.code_hash) # --------------------------------------------------------------------------- # POST /api/auth/totp/disable # --------------------------------------------------------------------------- def _do_full_enable(client: TestClient, csrf: str) -> str: """Perform setup + enable and return the secret.""" setup_body = _setup(client, csrf).json() secret = setup_body["secret"] _enable(client, csrf, pyotp.TOTP(secret).now()) return secret def test_disable_with_password_succeeds(totp_client): client, engine = totp_client csrf = _login(client) _do_full_enable(client, csrf) assert _status(client)["enabled"] is True resp = _disable(client, csrf, password="test-password") assert resp.status_code == 204, resp.text def test_disable_clears_enabled_flag(totp_client): client, engine = totp_client csrf = _login(client) _do_full_enable(client, csrf) _disable(client, csrf, password="test-password") assert _status(client)["enabled"] is False def test_disable_clears_secret_and_recovery_codes(totp_client): client, engine = totp_client csrf = _login(client) _do_full_enable(client, csrf) _disable(client, csrf, password="test-password") with Session(engine) as db: user = db.scalar(select(AuthUser).where(AuthUser.username == "admin")) assert user is not None assert user.totp_secret is None assert user.totp_enabled is False codes = db.execute(select(RecoveryCode).where(RecoveryCode.user_id == user.id)).scalars().all() assert len(codes) == 0 def test_disable_with_totp_code_succeeds(totp_client): client, engine = totp_client csrf = _login(client) secret = _do_full_enable(client, csrf) current_code = pyotp.TOTP(secret).now() resp = _disable(client, csrf, code=current_code) assert resp.status_code == 204, resp.text assert _status(client)["enabled"] is False def test_disable_with_wrong_password_fails(totp_client): client, engine = totp_client csrf = _login(client) _do_full_enable(client, csrf) resp = _disable(client, csrf, password="wrong-password") assert resp.status_code == 400 assert _status(client)["enabled"] is True def test_disable_with_no_credential_fails(totp_client): client, engine = totp_client csrf = _login(client) _do_full_enable(client, csrf) resp = _disable(client, csrf) # neither password nor code assert resp.status_code == 400 # --------------------------------------------------------------------------- # GET /api/auth/totp — status endpoint # --------------------------------------------------------------------------- def test_status_does_not_return_secret(totp_client): client, engine = totp_client csrf = _login(client) setup_body = _setup(client, csrf).json() secret = setup_body["secret"] _enable(client, csrf, pyotp.TOTP(secret).now()) status_resp = client.get("/api/auth/totp") assert status_resp.status_code == 200 body = status_resp.json() # Must contain ONLY enabled assert set(body.keys()) == {"enabled"} assert "secret" not in str(body) assert "recovery_codes" not in str(body) assert "otpauth" not in str(body) def test_status_initially_disabled(totp_client): client, _ = totp_client _login(client) body = _status(client) assert body["enabled"] is False def test_status_after_enable(totp_client): client, _ = totp_client csrf = _login(client) setup_body = _setup(client, csrf).json() _enable(client, csrf, pyotp.TOTP(setup_body["secret"]).now()) body = _status(client) assert body["enabled"] is True # --------------------------------------------------------------------------- # Security: unauthenticated / missing CSRF # --------------------------------------------------------------------------- def test_setup_unauthenticated_returns_401(client: TestClient): resp = client.post("/api/auth/totp/setup", headers={"X-CSRF-Token": "x"}) assert resp.status_code == 401 def test_setup_missing_csrf_returns_403(totp_client): client, _ = totp_client _login(client) # No X-CSRF-Token header resp = client.post("/api/auth/totp/setup") assert resp.status_code == 403 def test_enable_unauthenticated_returns_401(client: TestClient): resp = client.post( "/api/auth/totp/enable", json={"code": "123456"}, headers={"X-CSRF-Token": "x"}, ) assert resp.status_code == 401 def test_enable_missing_csrf_returns_403(totp_client): client, _ = totp_client _login(client) resp = client.post("/api/auth/totp/enable", json={"code": "123456"}) assert resp.status_code == 403 def test_disable_unauthenticated_returns_401(client: TestClient): resp = client.post( "/api/auth/totp/disable", json={"password": "test-password"}, headers={"X-CSRF-Token": "x"}, ) assert resp.status_code == 401 def test_disable_missing_csrf_returns_403(totp_client): client, _ = totp_client _login(client) resp = client.post("/api/auth/totp/disable", json={"password": "test-password"}) assert resp.status_code == 403 def test_status_unauthenticated_returns_401(client: TestClient): resp = client.get("/api/auth/totp") assert resp.status_code == 401 # --------------------------------------------------------------------------- # M4-T06 — verify_recovery_code service-level tests # --------------------------------------------------------------------------- def _full_setup_and_enable(client: TestClient) -> tuple[str, list[str]]: """Log in, run setup, enable TOTP, return (secret, plaintext_recovery_codes).""" csrf = _login(client) setup_body = _setup(client, csrf).json() secret = setup_body["secret"] recovery_codes = setup_body["recovery_codes"] _enable(client, csrf, pyotp.TOTP(secret).now()) return secret, recovery_codes class TestVerifyRecoveryCodeService: """Unit-level checks for totp_service.verify_recovery_code (used in T06 login).""" def test_verify_recovery_code_returns_true_on_valid_code(self, totp_client): """verify_recovery_code returns True for an unused code.""" from app.models.auth import AuthUser from sqlalchemy import select from app.services import totp as totp_svc client, engine = totp_client _, recovery_codes = _full_setup_and_enable(client) code = recovery_codes[0] with Session(engine) as db: user = db.scalar(select(AuthUser).where(AuthUser.username == "admin")) result = totp_svc.verify_recovery_code(db, user=user, code=code) assert result is True def test_verify_recovery_code_returns_false_on_used_code(self, totp_client): """verify_recovery_code returns False when the code was already consumed.""" from app.models.auth import AuthUser from sqlalchemy import select from app.services import totp as totp_svc client, engine = totp_client _, recovery_codes = _full_setup_and_enable(client) code = recovery_codes[0] with Session(engine) as db: user = db.scalar(select(AuthUser).where(AuthUser.username == "admin")) # First call consumes assert totp_svc.verify_recovery_code(db, user=user, code=code) is True # Reload user after commit db.expire_all() user = db.scalar(select(AuthUser).where(AuthUser.username == "admin")) # Second call must fail assert totp_svc.verify_recovery_code(db, user=user, code=code) is False def test_verify_recovery_code_returns_false_on_invalid_code(self, totp_client): """verify_recovery_code returns False for a completely wrong code.""" from app.models.auth import AuthUser from sqlalchemy import select from app.services import totp as totp_svc client, engine = totp_client _full_setup_and_enable(client) with Session(engine) as db: user = db.scalar(select(AuthUser).where(AuthUser.username == "admin")) result = totp_svc.verify_recovery_code(db, user=user, code="xxxx-xxxx") assert result is False def test_verify_totp_code_correct(self, totp_client): """verify_totp_code returns True for the current TOTP code.""" from app.models.auth import AuthUser from sqlalchemy import select from app.services import totp as totp_svc client, engine = totp_client secret, _ = _full_setup_and_enable(client) with Session(engine) as db: user = db.scalar(select(AuthUser).where(AuthUser.username == "admin")) result = totp_svc.verify_totp_code(user, pyotp.TOTP(secret).now()) assert result is True def test_verify_totp_code_wrong(self, totp_client): """verify_totp_code returns False for a wrong code.""" from app.models.auth import AuthUser from sqlalchemy import select from app.services import totp as totp_svc client, engine = totp_client _full_setup_and_enable(client) with Session(engine) as db: user = db.scalar(select(AuthUser).where(AuthUser.username == "admin")) result = totp_svc.verify_totp_code(user, "000000") assert result is False def test_verify_totp_code_no_secret(self, totp_client): """verify_totp_code returns False if user has no secret (TOTP not set up).""" from app.models.auth import AuthUser from sqlalchemy import select from app.services import totp as totp_svc client, engine = totp_client with Session(engine) as db: user = db.scalar(select(AuthUser).where(AuthUser.username == "admin")) assert user.totp_secret is None result = totp_svc.verify_totp_code(user, "123456") assert result is False