Newer
Older
gnexus-creds / gnexus_creds / crypto.py
"""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()