import base64 import os import uuid import time from datetime import datetime, timedelta from functools import wraps import jwt from cryptography.hazmat.primitives.ciphers.aead import AESGCM from flask import current_app, request, g, jsonify from argon2 import PasswordHasher from argon2.exceptions import VerifyMismatchError, VerificationError, InvalidHashError def hash_auth_token(auth_hash: str) -> str: """Hash the client-derived PBKDF2 auth_hash with Argon2id before storing.""" ph = PasswordHasher( time_cost=current_app.config['ARGON2_TIME_COST'], memory_cost=current_app.config['ARGON2_MEMORY_COST'], parallelism=current_app.config['ARGON2_PARALLELISM'], ) return ph.hash(auth_hash) def verify_auth_token(auth_hash: str, stored_hash: str) -> bool: ph = PasswordHasher() try: return ph.verify(stored_hash, auth_hash) except (VerifyMismatchError, VerificationError, InvalidHashError): return False def _get_totp_key() -> bytes: """ Return the 32-byte AES key used for server-side TOTP secret encryption. The key is stored as a 64-char hex string in TOTP_ENCRYPTION_KEY config. """ hex_key = current_app.config.get('TOTP_ENCRYPTION_KEY', '') if not hex_key or len(hex_key) != 64: raise RuntimeError( 'TOTP_ENCRYPTION_KEY must be set to a 64-character hex string (32 bytes). ' 'Generate one with: python -c "import secrets; print(secrets.token_hex(32))"' ) return bytes.fromhex(hex_key) def encrypt_totp_secret(plaintext_secret: str) -> tuple[str, str]: """ Encrypt a plaintext TOTP base32 secret with AES-256-GCM. Returns (ciphertext_b64, iv_b64). """ key = _get_totp_key() iv = os.urandom(12) aesgcm = AESGCM(key) ciphertext = aesgcm.encrypt(iv, plaintext_secret.encode(), None) return base64.b64encode(ciphertext).decode(), base64.b64encode(iv).decode() def decrypt_totp_secret(ciphertext_b64: str, iv_b64: str) -> str: """ Decrypt a base64-encoded AES-256-GCM TOTP secret ciphertext. Returns the plaintext base32 secret string. """ key = _get_totp_key() iv = base64.b64decode(iv_b64) ciphertext = base64.b64decode(ciphertext_b64) aesgcm = AESGCM(key) return aesgcm.decrypt(iv, ciphertext, None).decode() def generate_tokens(user_id: int) -> dict: """Return access_token and refresh_token JWTs, each with a unique jti.""" now = datetime.utcnow() secret = current_app.config['JWT_SECRET_KEY'] access_payload = { 'sub': str(user_id), 'type': 'access', 'jti': str(uuid.uuid4()), 'iat': now, 'exp': now + current_app.config['JWT_ACCESS_TOKEN_EXPIRES'], } refresh_payload = { 'sub': str(user_id), 'type': 'refresh', 'jti': str(uuid.uuid4()), 'iat': now, 'exp': now + current_app.config['JWT_REFRESH_TOKEN_EXPIRES'], } return { 'access_token': jwt.encode(access_payload, secret, algorithm='HS256'), 'refresh_token': jwt.encode(refresh_payload, secret, algorithm='HS256'), } def generate_mfa_token(user_id: int) -> str: """Short-lived (5-min) single-use token issued after password but before TOTP.""" now = datetime.utcnow() payload = { 'sub': str(user_id), 'type': 'mfa', 'jti': str(uuid.uuid4()), 'iat': now, 'exp': now + timedelta(minutes=5), } return jwt.encode(payload, current_app.config['JWT_SECRET_KEY'], algorithm='HS256') def decode_token(token: str, expected_type: str = 'access', check_blacklist: bool = True) -> dict: """Decode and validate a JWT. Raises jwt.PyJWTError on any failure.""" secret = current_app.config['JWT_SECRET_KEY'] payload = jwt.decode(token, secret, algorithms=['HS256']) if payload.get('type') != expected_type: raise jwt.InvalidTokenError('Wrong token type') if check_blacklist: from app.models.token_blacklist import TokenBlacklist jti = payload.get('jti') if jti and TokenBlacklist.is_blacklisted(jti): raise jwt.InvalidTokenError('Token has been revoked') return payload def blacklist_token(token: str, token_type: str) -> None: """Add a JWT's jti to the blacklist. Silently ignores invalid tokens.""" try: payload = decode_token(token, expected_type=token_type, check_blacklist=False) jti = payload.get('jti') if not jti: return exp = payload.get('exp') expires_at = datetime.utcfromtimestamp(exp) if exp else datetime.utcnow() + timedelta(days=7) from app.models.token_blacklist import TokenBlacklist from app import db # Avoid duplicate if already blacklisted if not TokenBlacklist.query.filter_by(jti=jti).first(): entry = TokenBlacklist( jti=jti, user_id=int(payload.get('sub', 0)), expires_at=expires_at, ) db.session.add(entry) db.session.commit() # Opportunistic cleanup — runs in same transaction context TokenBlacklist.cleanup_expired() except Exception: pass # Never let blacklisting errors break the logout flow def require_jwt(f): """Decorator: validates Bearer token and sets g.current_user_id.""" @wraps(f) def decorated(*args, **kwargs): auth_header = request.headers.get('Authorization', '') if not auth_header.startswith('Bearer '): return jsonify({'error': 'Missing or invalid Authorization header'}), 401 token = auth_header[7:] try: payload = decode_token(token, expected_type='access') except jwt.ExpiredSignatureError: return jsonify({'error': 'Token expired'}), 401 except jwt.PyJWTError: return jsonify({'error': 'Invalid token'}), 401 g.current_user_id = int(payload['sub']) return f(*args, **kwargs) return decorated