Files
PassKeeper/app/services/auth_service.py
T

165 lines
5.9 KiB
Python

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