import base64 import functools import hashlib import hmac import json import logging import os from typing import Optional from cryptography.hazmat.backends import default_backend from cryptography.hazmat.primitives import hashes from cryptography.hazmat.primitives.ciphers import algorithms, Cipher, modes from cryptography.hazmat.primitives.ciphers.aead import AESGCM from cryptography.hazmat.primitives.kdf.hkdf import HKDF from cryptography.hazmat.primitives.kdf.pbkdf2 import PBKDF2HMAC from docsgpt.core.settings import settings logger = logging.getLogger(__name__) def _derive_key(user_id: str, salt: bytes) -> bytes: app_secret = settings.ENCRYPTION_SECRET_KEY password = f"{app_secret}#{user_id}".encode() kdf = PBKDF2HMAC( algorithm=hashes.SHA256(), length=32, salt=salt, iterations=100000, backend=default_backend(), ) return kdf.derive(password) def encrypt_credentials(credentials: dict, user_id: str) -> str: if not credentials: return "" try: salt = os.urandom(16) iv = os.urandom(16) key = _derive_key(user_id, salt) json_str = json.dumps(credentials) cipher = Cipher(algorithms.AES(key), modes.CBC(iv), backend=default_backend()) encryptor = cipher.encryptor() padded_data = _pad_data(json_str.encode()) encrypted_data = encryptor.update(padded_data) + encryptor.finalize() result = salt + iv + encrypted_data return base64.b64encode(result).decode() except Exception as e: logger.warning(f"Failed to encrypt credentials: {e}") return "" def decrypt_credentials(encrypted_data: str, user_id: str) -> dict: if not encrypted_data: return {} try: data = base64.b64decode(encrypted_data.encode()) salt = data[:16] iv = data[16:32] encrypted_content = data[32:] key = _derive_key(user_id, salt) cipher = Cipher(algorithms.AES(key), modes.CBC(iv), backend=default_backend()) decryptor = cipher.decryptor() decrypted_padded = decryptor.update(encrypted_content) + decryptor.finalize() decrypted_data = _unpad_data(decrypted_padded) return json.loads(decrypted_data.decode()) except Exception as e: logger.warning(f"Failed to decrypt credentials: {e}") return {} def _pad_data(data: bytes) -> bytes: block_size = 16 padding_len = block_size - (len(data) % block_size) padding = bytes([padding_len]) * padding_len return data + padding def _unpad_data(data: bytes) -> bytes: padding_len = data[-1] return data[:-padding_len] # --------------------------------------------------------------------------- # Envelope v2: connection credentials # --------------------------------------------------------------------------- # # ``v2::`` # # AES-256-GCM, so a tampered blob fails to decrypt instead of returning # garbage. The master key is derived once per process from # ENCRYPTION_SECRET_KEY (PBKDF2, cached); each record gets its own key from # HKDF(master, salt, owner id), which keeps the v1 owner binding without # paying 100k PBKDF2 iterations on every token read in the worker. The owner # id is also the GCM associated data, so a blob copied onto another user's # row does not decrypt. ``key_id`` names the master key, so a blob written # under ENCRYPTION_SECRET_KEY_PREVIOUS is still readable during a rotation. _V2_PREFIX = "v2" _V2_MASTER_SALT = b"docsgpt-credentials-v2" _V2_ITERATIONS = 200_000 _V2_SALT_BYTES = 16 _V2_NONCE_BYTES = 12 DEFAULT_ENCRYPTION_KEY = "default-docsgpt-encryption-key" class CredentialDecryptionError(Exception): """A stored credential could not be decrypted (wrong key, tampering, bad format).""" @functools.lru_cache(maxsize=8) def _master_key(secret: str) -> bytes: kdf = PBKDF2HMAC( algorithm=hashes.SHA256(), length=32, salt=_V2_MASTER_SALT, iterations=_V2_ITERATIONS, backend=default_backend(), ) return kdf.derive(secret.encode()) def _key_id(master: bytes) -> str: return hmac.new(master, b"docsgpt-key-id", hashlib.sha256).hexdigest()[:8] def _record_key(master: bytes, owner_id: str, salt: bytes) -> bytes: return HKDF( algorithm=hashes.SHA256(), length=32, salt=salt, info=b"docsgpt-v2|" + owner_id.encode(), backend=default_backend(), ).derive(master) def _candidate_keys() -> dict[str, bytes]: """Master keys this process can decrypt with, by key id (current first).""" keys: dict[str, bytes] = {} for secret in (settings.ENCRYPTION_SECRET_KEY, settings.ENCRYPTION_SECRET_KEY_PREVIOUS): if secret: master = _master_key(secret) keys.setdefault(_key_id(master), master) return keys def current_key_id() -> str: """Key id of ENCRYPTION_SECRET_KEY, as written into new v2 blobs.""" return _key_id(_master_key(settings.ENCRYPTION_SECRET_KEY)) def is_default_encryption_key() -> bool: """Whether ENCRYPTION_SECRET_KEY is still the public default.""" return settings.ENCRYPTION_SECRET_KEY == DEFAULT_ENCRYPTION_KEY def encrypt_json(data: dict, owner_id: str) -> str: """Encrypt ``data`` for ``owner_id`` into a v2 envelope. Args: data: JSON-serialisable credentials. owner_id: The user the credentials belong to; decryption needs it. Returns: The ``v2::`` string. """ master = _master_key(settings.ENCRYPTION_SECRET_KEY) key_id = _key_id(master) salt = os.urandom(_V2_SALT_BYTES) nonce = os.urandom(_V2_NONCE_BYTES) key = _record_key(master, owner_id, salt) plaintext = json.dumps(data, separators=(",", ":")).encode() ciphertext = AESGCM(key).encrypt(nonce, plaintext, owner_id.encode()) payload = base64.b64encode(salt + nonce + ciphertext).decode() return f"{_V2_PREFIX}:{key_id}:{payload}" def envelope_key_id(blob: str) -> Optional[str]: """The key id a v2 blob was written with, or None for anything else.""" parts = (blob or "").split(":", 2) if len(parts) != 3 or parts[0] != _V2_PREFIX: return None return parts[1] def decrypt_json(blob: str, owner_id: str) -> dict: """Decrypt a v2 envelope written for ``owner_id``. Raises: CredentialDecryptionError: The blob is malformed, was written with a key this process does not have, belongs to another owner, or was tampered with. """ key_id = envelope_key_id(blob) if key_id is None: raise CredentialDecryptionError("Not a v2 credential envelope") master = _candidate_keys().get(key_id) if master is None: raise CredentialDecryptionError("Credential was encrypted with an unknown key") try: raw = base64.b64decode(blob.split(":", 2)[2].encode(), validate=True) salt = raw[:_V2_SALT_BYTES] nonce = raw[_V2_SALT_BYTES:_V2_SALT_BYTES + _V2_NONCE_BYTES] ciphertext = raw[_V2_SALT_BYTES + _V2_NONCE_BYTES:] key = _record_key(master, owner_id, salt) plaintext = AESGCM(key).decrypt(nonce, ciphertext, owner_id.encode()) data = json.loads(plaintext.decode()) except Exception as exc: raise CredentialDecryptionError("Credential could not be decrypted") from exc if not isinstance(data, dict): raise CredentialDecryptionError("Credential payload is not an object") return data