"""Encryption helpers."""
import base64
import hashlib
import os
import secrets
from functools import lru_cache
from cryptography.exceptions import InvalidTag
from cryptography.hazmat.primitives.ciphers.aead import AESGCM
ALGORITHM = "AESGCM"
def _b64e(value: bytes) -> str:
return base64.urlsafe_b64encode(value).decode().rstrip("=")
def _b64d(value: str) -> bytes:
padding = "=" * (-len(value) % 4)
return base64.urlsafe_b64decode(value + padding)
def _derive_with_salt(master_key: str, salt: str) -> bytes:
# n=2**15, r=8 needs 128*n*r = 32MB; raise maxmem above OpenSSL's
# built-in default limit, or scrypt raises "memory limit exceeded"
return hashlib.scrypt(
master_key.encode(), salt=salt.encode(), n=2**15, r=8, p=1, dklen=32, maxmem=64 * 2**20
)
def _derive_legacy(master_key: str) -> bytes:
return hashlib.sha256(master_key.encode()).digest()
@lru_cache(maxsize=8)
def _derive_cached(master_key: str, salt: str) -> bytes:
# scrypt costs ~100ms; without the cache every request re-pays it
if salt:
return _derive_with_salt(master_key, salt)
return _derive_legacy(master_key)
def derive_master_key(master_key: str, salt: str | None = None) -> bytes:
"""Master key from the master string and the installation's salt.
With a salt this is key-stretched (scrypt); without one it falls back to
the legacy plain sha256 so pre-salt installs keep working."""
return _derive_cached(master_key, salt or "")
def new_raw_key() -> bytes:
return AESGCM.generate_key(bit_length=256)
def encrypt_bytes(key: bytes, plaintext: bytes, *, aad: bytes | None = None) -> tuple[bytes, bytes]:
nonce = os.urandom(12)
ciphertext = AESGCM(key).encrypt(nonce, plaintext, aad)
return nonce, ciphertext
def decrypt_bytes(
key: bytes, nonce: bytes, ciphertext: bytes, *, aad: bytes | None = None
) -> bytes:
return AESGCM(key).decrypt(nonce, ciphertext, aad)
def encrypt_text(key: bytes, key_id: str, value: str, *, aad: str | None = None) -> dict:
nonce, ciphertext = encrypt_bytes(key, value.encode(), aad=aad.encode() if aad else None)
return {
"algorithm": ALGORITHM,
"key_id": key_id,
"nonce": _b64e(nonce),
"ciphertext": _b64e(ciphertext),
}
def decrypt_text(key: bytes, payload: dict, *, aad: str | None = None) -> str:
ciphertext = _b64d(payload["ciphertext"])
nonce = _b64d(payload["nonce"])
if aad:
try:
return decrypt_bytes(key, nonce, ciphertext, aad=aad.encode()).decode()
except InvalidTag:
# pre-AAD rows (written before the context binding) carry no aad —
# try the unbound form before giving up; anything still failing is
# tampering or a corrupt row and raises InvalidTag below
return decrypt_bytes(key, nonce, ciphertext).decode()
return decrypt_bytes(key, nonce, ciphertext).decode()
def token_secret() -> str:
return secrets.token_urlsafe(32)
def token_hash(token: str) -> str:
return hashlib.sha256(token.encode()).hexdigest()