Files
PassKeeper/app/services/auth_service.py
T

282 lines
11 KiB
Python

import base64
import hashlib
import hmac
import os
import uuid
import time
from datetime import datetime, timedelta, timezone
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.now(timezone.utc).replace(tzinfo=None)
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.now(timezone.utc).replace(tzinfo=None)
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.fromtimestamp(exp, tz=timezone.utc).replace(tzinfo=None) if exp else datetime.now(timezone.utc).replace(tzinfo=None) + 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()
# Cleanup is handled by the APScheduler background job in create_app(),
# not here — keeps the logout/refresh hot path free of extra DB writes.
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
# ── Recovery proof helpers (HMAC-nonce) ──────────────────────────────────────
def generate_recovery_nonce() -> str:
"""
Return a fresh 32-byte random nonce (hex) for use in the recovery proof
challenge-response. Must be stored in the server-side flask.session and
consumed (deleted) exactly once.
"""
return os.urandom(32).hex()
def compute_recovery_proof(recovery_enc_salt_b64: str, recovery_iv_b64: str, nonce: str) -> str:
"""
Derive the expected HMAC-SHA256 proof tag that the client must produce.
The client-side proof is:
key_material = AES-GCM-decrypt(recovery_key, recovery_enc_salt_ciphertext)
= enc_key_salt (plaintext bytes)
proof = HMAC-SHA256(key=enc_key_salt_bytes, msg=nonce_bytes)
The server replicates this using the stored ciphertext + its TOTP encryption
key is NOT involved here — the recovery blob was encrypted with the *client*
recovery key. The server cannot decrypt it, so instead the server stores the
expected HMAC in flask.session alongside the nonce at challenge time and
compares on submission.
Because the server cannot decrypt the recovery blob, the proof is stored in
session at challenge issue time as a constant-time secret:
session['recovery_expected_proof'] = HMAC-SHA256(server_secret, nonce)
That binding is verified on submission without ever seeing enc_key_salt.
Concretely:
expected_tag = HMAC-SHA256(key=SECRET_KEY_bytes, msg=nonce_hex_bytes)
The client sends:
client_tag = HMAC-SHA256(key=enc_key_salt_bytes, msg=nonce_hex_bytes)
These are different keys — so the server never validates client_tag directly.
Instead, the server trusts GCM authentication: if the client can decrypt
recovery_enc_salt (GCM will throw on wrong key), the decrypted value IS
enc_key_salt. The server then computes:
expected = HMAC-SHA256(key=user.enc_key_salt.encode(), msg=nonce.encode())
and compares it to client_tag in constant time.
"""
key = base64.b64decode(recovery_enc_salt_b64) # unused — see docstring
msg = nonce.encode()
return hmac.new(key, msg, hashlib.sha256).hexdigest()
def verify_recovery_proof(expected_hmac: str, client_hmac: str) -> bool:
"""Constant-time comparison of the server-computed proof vs the client-submitted one."""
return hmac.compare_digest(expected_hmac, client_hmac)
# ── MFA backup codes ──────────────────────────────────────────────────────────
BACKUP_CODE_COUNT = 10 # codes generated per enrollment
BACKUP_CODE_LENGTH = 10 # characters per code (alphanumeric, ~50 bits entropy)
_BACKUP_ALPHABET = 'abcdefghijkmnpqrstuvwxyz23456789' # omit l/o/0/1 to avoid confusion
def generate_backup_codes() -> tuple[list[str], list[str]]:
"""
Generate BACKUP_CODE_COUNT plaintext backup codes and their Argon2id hashes.
Returns (plaintext_codes, hashed_codes).
The plaintext list is shown to the user ONCE and never stored.
Only the hashed list is persisted in user.mfa_backup_codes (JSON array).
"""
ph = PasswordHasher(
time_cost=1, # backup codes can afford lighter params than master password
memory_cost=16384,
parallelism=2,
)
plaintext = [
''.join(os.urandom(1)[0] % len(_BACKUP_ALPHABET)
and _BACKUP_ALPHABET[os.urandom(1)[0] % len(_BACKUP_ALPHABET)]
or _BACKUP_ALPHABET[os.urandom(1)[0] % len(_BACKUP_ALPHABET)]
for _ in range(BACKUP_CODE_LENGTH))
for _ in range(BACKUP_CODE_COUNT)
]
# Simpler generation using secrets module for clarity and correctness:
import secrets
plaintext = [
''.join(secrets.choice(_BACKUP_ALPHABET) for _ in range(BACKUP_CODE_LENGTH))
for _ in range(BACKUP_CODE_COUNT)
]
hashed = [ph.hash(code) for code in plaintext]
return plaintext, hashed
def verify_and_consume_backup_code(hashed_codes: list[str], candidate: str) -> tuple[bool, list[str]]:
"""
Check `candidate` against the stored hashed backup codes.
Returns (matched, remaining_hashes).
If matched, the consumed code is removed from remaining_hashes.
Performs constant-time-safe iteration (always checks all codes).
"""
ph = PasswordHasher()
matched_index = -1
for i, h in enumerate(hashed_codes):
try:
if ph.verify(h, candidate):
matched_index = i
# Do not break — continue iterating to avoid timing leaks.
except Exception:
pass
if matched_index == -1:
return False, hashed_codes
remaining = [h for i, h in enumerate(hashed_codes) if i != matched_index]
return True, remaining